Skip to content

Commit 03f5ffa

Browse files
authored
impl(bigquery): add query polling and complete query (googleapis#5896)
1 parent be14699 commit 03f5ffa

3 files changed

Lines changed: 322 additions & 3 deletions

File tree

src/bigquery/Cargo.toml

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -34,12 +34,13 @@ google-cloud-bigquery-v2 = { workspace = true }
3434
thiserror.workspace = true
3535
wkt.workspace = true
3636
gaxi = { workspace = true, features = ["_internal-common", "_internal-grpc-client", "_internal-http-client"] }
37+
tokio = { workspace = true, features = ["time"] }
3738

3839
[dev-dependencies]
3940
anyhow.workspace = true
4041
mockall.workspace = true
4142
test-case.workspace = true
42-
tokio = { workspace = true, features = ["macros", "rt-multi-thread"] }
43+
tokio = { workspace = true, features = ["macros", "rt-multi-thread", "test-util"] }
4344

4445
[features]
4546
default = ["default-rustls-provider"]

src/bigquery/src/query.rs

Lines changed: 23 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -36,8 +36,13 @@ pub type Result<T> = std::result::Result<T, crate::error::QueryError>;
3636
pub(crate) mod tests {
3737
use google_cloud_bigquery_v2::Result;
3838
use google_cloud_bigquery_v2::client::JobService;
39-
use google_cloud_bigquery_v2::model::{InsertJobRequest, Job, PostQueryRequest, QueryResponse};
39+
use google_cloud_bigquery_v2::model::{
40+
GetQueryResultsRequest, GetQueryResultsResponse, InsertJobRequest, Job, PostQueryRequest,
41+
QueryResponse,
42+
};
4043
use google_cloud_gax::options::RequestOptions;
44+
use google_cloud_gax::polling_backoff_policy::PollingBackoffPolicy;
45+
use google_cloud_gax::polling_state::PollingState;
4146
use google_cloud_gax::response::Response;
4247
use std::sync::Arc;
4348

@@ -55,10 +60,27 @@ pub(crate) mod tests {
5560
req: PostQueryRequest,
5661
options: RequestOptions,
5762
) -> Result<Response<QueryResponse>>;
63+
async fn get_query_results(
64+
&self,
65+
req: GetQueryResultsRequest,
66+
options: RequestOptions,
67+
) -> Result<Response<GetQueryResultsResponse>>;
68+
}
69+
}
70+
71+
mockall::mock! {
72+
#[derive(Debug)]
73+
pub BackoffPolicy {}
74+
impl PollingBackoffPolicy for BackoffPolicy {
75+
fn wait_period(&self, _state: &PollingState) -> std::time::Duration;
5876
}
5977
}
6078

6179
pub(crate) fn create_job_service(mock: MockJobService) -> Arc<JobService> {
6280
Arc::new(JobService::from_stub::<MockJobService>(Arc::new(mock)))
6381
}
82+
83+
pub(crate) fn create_test_backoff_policy() -> MockBackoffPolicy {
84+
MockBackoffPolicy::new()
85+
}
6486
}

src/bigquery/src/query/query_handle.rs

Lines changed: 297 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -12,8 +12,17 @@
1212
// See the License for the specific language governing permissions and
1313
// limitations under the License.
1414

15+
use crate::error::QueryError;
16+
use crate::query::{Result, Schema};
1517
use google_cloud_bigquery_v2::client::JobService;
16-
use google_cloud_bigquery_v2::model::{Job, JobReference, QueryResponse};
18+
use google_cloud_bigquery_v2::model::{
19+
GetQueryResultsRequest, GetQueryResultsResponse, Job, JobReference, QueryResponse,
20+
};
21+
use google_cloud_gax::backoff_policy::BackoffPolicy;
22+
use google_cloud_gax::exponential_backoff::ExponentialBackoffBuilder;
23+
use google_cloud_gax::polling_backoff_policy::PollingBackoffPolicy;
24+
use google_cloud_gax::polling_state::PollingState;
25+
use std::collections::VecDeque;
1726
use std::sync::Arc;
1827

1928
/// A handle representing a running query.
@@ -25,3 +34,290 @@ pub struct Query {
2534
pub(crate) initial_job: Option<Job>,
2635
pub(crate) initial_response: Option<QueryResponse>,
2736
}
37+
38+
impl Query {
39+
/// Periodically checks the status of the background job until it finishes.
40+
/// Returns an error if a remote service or connection failure happens during polling.
41+
pub async fn until_done(&self) -> Result<CompleteQuery> {
42+
if let (true, Some(initial_response)) = (self.completed, &self.initial_response) {
43+
return Ok(CompleteQuery::from_query_response(self, initial_response));
44+
}
45+
46+
let job_ref = self
47+
.job_ref
48+
.as_ref()
49+
.expect("query job should have job reference at this point");
50+
let backoff_policy = Arc::new(
51+
ExponentialBackoffBuilder::default()
52+
.with_initial_delay(std::time::Duration::from_secs(10))
53+
.build()
54+
.expect("valid backoff configuration"),
55+
);
56+
let res = poll_query_results(&self.job_service, job_ref, backoff_policy).await?;
57+
Ok(CompleteQuery::from_get_query_results_response(self, res))
58+
}
59+
}
60+
61+
/// A handle representing a successfully completed query ready for reading.
62+
#[derive(Debug, Clone)]
63+
pub struct CompleteQuery {
64+
pub(crate) job_service: Arc<JobService>,
65+
pub(crate) job_ref: Option<JobReference>,
66+
}
67+
68+
impl CompleteQuery {
69+
pub(crate) fn from_get_query_results_response(
70+
q: &Query,
71+
_res: GetQueryResultsResponse,
72+
) -> Self {
73+
// TODO(#5592): hold cached rows, page token, schema and query metadata here.
74+
Self {
75+
job_service: q.job_service.clone(),
76+
job_ref: q.job_ref.clone(),
77+
}
78+
}
79+
80+
pub(crate) fn from_query_response(q: &Query, _res: &QueryResponse) -> Self {
81+
// TODO(#5592): hold cached rows, page token, schema and query metadata here.
82+
Self {
83+
job_service: q.job_service.clone(),
84+
job_ref: q.job_ref.clone(),
85+
}
86+
}
87+
}
88+
89+
/// Helper function to poll getQueryResults until a job finishes.
90+
pub(crate) async fn poll_query_results(
91+
job_service: &JobService,
92+
job_ref: &JobReference,
93+
backoff_policy: Arc<dyn PollingBackoffPolicy>,
94+
) -> Result<GetQueryResultsResponse> {
95+
let mut state = PollingState::default();
96+
97+
loop {
98+
let mut req = GetQueryResultsRequest::new()
99+
.set_max_results(0u32)
100+
.set_project_id(job_ref.project_id.clone())
101+
.set_job_id(job_ref.job_id.clone());
102+
if let Some(location) = job_ref.location.clone() {
103+
req = req.set_location(location);
104+
}
105+
106+
let res = job_service
107+
.get_query_results()
108+
.with_request(req)
109+
.send()
110+
.await?;
111+
112+
if !res.errors.is_empty() {
113+
// TODO(#5592): handle jobBackendError and other transient/retryable errors.
114+
return Err(QueryError::JobFailed { errors: res.errors });
115+
}
116+
117+
let completed = res.job_complete.unwrap_or(false);
118+
if completed {
119+
return Ok(res);
120+
}
121+
122+
let delay = backoff_policy.wait_period(&state);
123+
tokio::time::sleep(delay).await;
124+
// TODO(#5592): limit retry attempts or add cancellation mechanism
125+
state.attempt_count += 1;
126+
}
127+
}
128+
129+
#[cfg(test)]
130+
mod tests {
131+
use std::time::Duration;
132+
133+
use super::*;
134+
use crate::query::tests::{
135+
MockBackoffPolicy, MockJobService, create_job_service, create_test_backoff_policy,
136+
};
137+
use google_cloud_bigquery_v2::model::{
138+
ErrorProto, GetQueryResultsResponse, JobReference, QueryResponse,
139+
};
140+
use google_cloud_gax::error::Error as GaxError;
141+
use google_cloud_gax::error::rpc::{Code, Status};
142+
use google_cloud_gax::response::Response;
143+
144+
use test_case::test_case;
145+
146+
type TestResult = anyhow::Result<()>;
147+
148+
#[tokio::test]
149+
async fn test_query_until_done_already_completed() -> TestResult {
150+
let job_service = create_job_service(MockJobService::new());
151+
let job_ref = JobReference::new()
152+
.set_project_id("some_project")
153+
.set_job_id("some_job_id");
154+
let query_res = QueryResponse::new()
155+
.set_job_complete(true)
156+
.set_job_reference(job_ref.clone());
157+
158+
let query = Query {
159+
job_service,
160+
job_ref: Some(job_ref),
161+
completed: true,
162+
initial_job: None,
163+
initial_response: Some(query_res),
164+
};
165+
166+
let completed = query.until_done().await?;
167+
assert_eq!(completed.job_ref.unwrap().job_id, "some_job_id");
168+
169+
Ok(())
170+
}
171+
172+
#[tokio::test]
173+
async fn test_query_until_done_polls_success() -> TestResult {
174+
let mut mock = MockJobService::new();
175+
mock.expect_get_query_results()
176+
.returning(|req, _| {
177+
assert_eq!(req.project_id, "some_project");
178+
assert_eq!(req.job_id, "some_job_id");
179+
assert_eq!(req.max_results, Some(0));
180+
assert_eq!(req.location, "us-central1");
181+
let res = GetQueryResultsResponse::new()
182+
.set_job_complete(true)
183+
.set_job_reference(JobReference::new().set_job_id(req.job_id));
184+
Ok(Response::from(res))
185+
})
186+
.times(1);
187+
let job_service = create_job_service(mock);
188+
let job_ref = JobReference::new()
189+
.set_project_id("some_project")
190+
.set_job_id("some_job_id")
191+
.set_location("us-central1");
192+
193+
let query = Query {
194+
job_service,
195+
job_ref: Some(job_ref),
196+
completed: false,
197+
initial_job: None,
198+
initial_response: None,
199+
};
200+
201+
let completed = query.until_done().await?;
202+
assert_eq!(completed.job_ref.unwrap().job_id, "some_job_id");
203+
204+
Ok(())
205+
}
206+
207+
#[tokio::test(start_paused = true)]
208+
async fn test_poll_query_results_loops_until_complete() -> TestResult {
209+
let mut mock = MockJobService::new();
210+
let mut backoff_policy = create_test_backoff_policy();
211+
backoff_policy
212+
.expect_wait_period()
213+
.times(2)
214+
.return_const(Duration::from_millis(1));
215+
216+
let mut seq = mockall::Sequence::new();
217+
218+
mock.expect_get_query_results()
219+
.in_sequence(&mut seq)
220+
.times(2)
221+
.returning(|_, _| {
222+
Ok(Response::from(
223+
GetQueryResultsResponse::new().set_job_complete(false),
224+
))
225+
});
226+
227+
mock.expect_get_query_results()
228+
.in_sequence(&mut seq)
229+
.times(1)
230+
.returning(|_, _| {
231+
Ok(Response::from(
232+
GetQueryResultsResponse::new().set_job_complete(true),
233+
))
234+
});
235+
236+
let job_service = create_job_service(mock);
237+
let job_ref = JobReference::new()
238+
.set_project_id("some_project")
239+
.set_job_id("some_job_id");
240+
241+
let res = poll_query_results(&job_service, &job_ref, Arc::new(backoff_policy)).await?;
242+
243+
assert!(res.job_complete.unwrap(), "{res:?}");
244+
245+
Ok(())
246+
}
247+
248+
#[tokio::test]
249+
async fn test_query_until_done_job_failed_error() -> TestResult {
250+
let mut mock = MockJobService::new();
251+
mock.expect_get_query_results().returning(|req, _| {
252+
assert_eq!(req.project_id, "some_project");
253+
assert_eq!(req.job_id, "some_job_id");
254+
assert_eq!(req.max_results, Some(0));
255+
let err_proto = ErrorProto::new()
256+
.set_reason("invalidQuery")
257+
.set_message("Syntax error");
258+
let res = GetQueryResultsResponse::new().set_errors(vec![err_proto]);
259+
Ok(Response::from(res))
260+
});
261+
let job_service = create_job_service(mock);
262+
let job_ref = JobReference::new()
263+
.set_project_id("some_project")
264+
.set_job_id("some_job_id");
265+
266+
let query = Query {
267+
job_service,
268+
job_ref: Some(job_ref),
269+
completed: false,
270+
initial_job: None,
271+
initial_response: None,
272+
};
273+
274+
let err = query.until_done().await.unwrap_err();
275+
let errors = match err {
276+
QueryError::JobFailed { errors } => errors,
277+
_ => panic!("expected QueryError::JobFailed, got {err:?}"),
278+
};
279+
assert_eq!(
280+
errors,
281+
[ErrorProto::new()
282+
.set_reason("invalidQuery")
283+
.set_message("Syntax error")]
284+
);
285+
286+
Ok(())
287+
}
288+
289+
#[tokio::test]
290+
async fn test_query_until_done_rpc_error() -> TestResult {
291+
let mut mock = MockJobService::new();
292+
mock.expect_get_query_results().returning(|req, _| {
293+
assert_eq!(req.project_id, "some_project");
294+
assert_eq!(req.job_id, "some_job_id");
295+
assert_eq!(req.max_results, Some(0));
296+
let status = Status::default()
297+
.set_code(Code::InvalidArgument)
298+
.set_message("simulated bad request");
299+
Err(GaxError::service(status))
300+
});
301+
let job_service = create_job_service(mock);
302+
let job_ref = JobReference::new()
303+
.set_project_id("some_project")
304+
.set_job_id("some_job_id");
305+
306+
let query = Query {
307+
job_service,
308+
job_ref: Some(job_ref),
309+
completed: false,
310+
initial_job: None,
311+
initial_response: None,
312+
};
313+
314+
let err = query.until_done().await.unwrap_err();
315+
let source = match err {
316+
QueryError::Rpc { source } => source,
317+
_ => panic!("expected QueryError::Rpc, got {err:?}"),
318+
};
319+
assert_eq!(source.status().unwrap().code, Code::InvalidArgument);
320+
321+
Ok(())
322+
}
323+
}

0 commit comments

Comments
 (0)