Skip to repository content2110 lines · 73.3 KB · rust
tenant.openagents/omega
No repository description is available.
OpenAgents Git authority 2026-07-28T03:53:51.044Z Public web read
NIP-34 coordinate
30617:7649603503856e5148d571eac2766b288a8ff1e9e35d380337a1d2b0015b4f92:omegaMaintainersHidden in public view
References2 branches · 1 tag
Read-only clone
git clone https://openagents.com/git/tenant.openagents/omega.gitBrowse files
pull_examples.rs
1use anyhow::{Context as _, Result};
2use flate2::read::GzDecoder;
3use gpui::BackgroundExecutor;
4use http_client::{AsyncBody, HttpClient, Method, Request};
5use indoc::indoc;
6use serde::Deserialize;
7use serde_json::{Value as JsonValue, json};
8use std::collections::HashMap;
9use std::fmt::Write as _;
10use std::io::Read;
11use std::sync::Arc;
12use std::time::Duration;
13use telemetry_events::EditPredictionRating;
14
15use zeta_prompt::{Zeta2PromptInput, ZetaFormat, excerpt_range_for_format};
16
17use crate::PredictionProvider;
18use crate::example::{Example, ExamplePrediction, ExamplePrompt};
19use crate::progress::{InfoStyle, Progress, Step};
20use edit_prediction::example_spec::{ExampleSpec, TelemetrySource};
21
22pub(crate) const SNOWFLAKE_SUCCESS_CODE: &str = "090001";
23pub(crate) const SNOWFLAKE_ASYNC_IN_PROGRESS_CODE: &str = "333334";
24const SNOWFLAKE_TIMEOUT_CODE: &str = "000630";
25
26/// Minimum Zed version for filtering captured examples.
27/// For example, `MinCaptureVersion { major: 0, minor: 224, patch: 1 }` means only pull
28/// examples where `zed_version >= 0.224.1`. The `major` component is required because Zed
29/// moved from the `0.<minor>.<patch>` scheme to `1.<minor>.<patch>`; comparing on `minor`
30/// alone would exclude all `1.*` versions (whose `minor` resets to small values).
31#[derive(Clone, Copy, Debug)]
32pub struct MinCaptureVersion {
33 pub major: u32,
34 pub minor: u32,
35 pub patch: u32,
36}
37
38pub(crate) const POLL_INTERVAL: Duration = Duration::from_secs(2);
39const PARTITION_FETCH_MAX_RETRIES: usize = 3;
40const PARTITION_FETCH_RETRY_DELAYS: [Duration; PARTITION_FETCH_MAX_RETRIES] = [
41 Duration::from_millis(500),
42 Duration::from_secs(1),
43 Duration::from_secs(2),
44];
45
46/// Parse an input token of the form `captured-after:{timestamp}`.
47pub fn parse_captured_after_input(input: &str) -> Option<&str> {
48 input.strip_prefix("captured-after:")
49}
50
51/// Parse an input token of the form `accepted-after:{timestamp}`.
52pub fn parse_accepted_after_input(input: &str) -> Option<&str> {
53 input.strip_prefix("accepted-after:")
54}
55
56/// Parse an input token of the form `rejected-after:{timestamp}`.
57pub fn parse_rejected_after_input(input: &str) -> Option<(bool, &str)> {
58 if let Some(timestamp) = input.strip_prefix("rejected-after:") {
59 Some((false, timestamp))
60 } else if let Some(timestamp) = input.strip_prefix("explicitly-rejected-after:") {
61 Some((true, timestamp))
62 } else {
63 None
64 }
65}
66
67/// Parse an input token of the form `requested-after:{timestamp}`.
68pub fn parse_requested_after_input(input: &str) -> Option<&str> {
69 input.strip_prefix("requested-after:")
70}
71
72/// Parse an input token of the form `settled-after:{timestamp}`.
73pub fn parse_settled_after_input(input: &str) -> Option<&str> {
74 input.strip_prefix("settled-after:")
75}
76
77/// Parse an input token of the form `rated-after:{timestamp}`, `rated-positive-after:{timestamp}`,
78/// or `rated-negative-after:{timestamp}`.
79/// Returns `(timestamp, Option<EditPredictionRating>)` where `None` means all ratings.
80pub fn parse_rated_after_input(input: &str) -> Option<(&str, Option<EditPredictionRating>)> {
81 if let Some(timestamp) = input.strip_prefix("rated-positive-after:") {
82 Some((timestamp, Some(EditPredictionRating::Positive)))
83 } else if let Some(timestamp) = input.strip_prefix("rated-negative-after:") {
84 Some((timestamp, Some(EditPredictionRating::Negative)))
85 } else if let Some(timestamp) = input.strip_prefix("rated-after:") {
86 Some((timestamp, None))
87 } else {
88 None
89 }
90}
91
92#[derive(Debug, Clone, Deserialize)]
93#[serde(rename_all = "camelCase")]
94pub(crate) struct SnowflakeStatementResponse {
95 #[serde(default)]
96 pub(crate) data: Vec<Vec<JsonValue>>,
97 #[serde(default)]
98 pub(crate) result_set_meta_data: Option<SnowflakeResultSetMetaData>,
99 #[serde(default)]
100 pub(crate) code: Option<String>,
101 #[serde(default)]
102 pub(crate) message: Option<String>,
103 #[serde(default)]
104 pub(crate) statement_handle: Option<String>,
105}
106
107#[derive(Debug, Clone, Deserialize)]
108#[serde(rename_all = "camelCase")]
109pub(crate) struct SnowflakeResultSetMetaData {
110 #[serde(default, rename = "rowType")]
111 row_type: Vec<SnowflakeColumnMeta>,
112 #[serde(default)]
113 num_rows: Option<i64>,
114 #[serde(default)]
115 partition_info: Vec<SnowflakePartitionInfo>,
116}
117
118#[derive(Debug, Clone, Deserialize)]
119#[serde(rename_all = "camelCase")]
120struct SnowflakePartitionInfo {}
121
122#[derive(Debug, Clone, Deserialize)]
123struct SnowflakeColumnMeta {
124 #[serde(default)]
125 name: String,
126}
127
128async fn run_sql_with_polling(
129 http_client: Arc<dyn HttpClient>,
130 base_url: &str,
131 token: &str,
132 request: &serde_json::Value,
133 step_progress: &crate::progress::StepProgress,
134 background_executor: BackgroundExecutor,
135) -> Result<SnowflakeStatementResponse> {
136 let mut response = run_sql(http_client.clone(), base_url, token, request).await?;
137
138 if response.code.as_deref() == Some(SNOWFLAKE_ASYNC_IN_PROGRESS_CODE) {
139 let statement_handle = response
140 .statement_handle
141 .as_ref()
142 .context("async query response missing statementHandle")?
143 .clone();
144
145 for attempt in 0.. {
146 step_progress.set_substatus(format!("polling ({attempt})"));
147
148 background_executor.timer(POLL_INTERVAL).await;
149
150 response = fetch_partition_with_retries(
151 http_client.clone(),
152 base_url,
153 token,
154 &statement_handle,
155 0,
156 background_executor.clone(),
157 )
158 .await?;
159
160 if response.code.as_deref() != Some(SNOWFLAKE_ASYNC_IN_PROGRESS_CODE) {
161 break;
162 }
163 }
164 }
165
166 Ok(response)
167}
168
169struct SnowflakeConfig {
170 token: String,
171 base_url: String,
172 role: Option<String>,
173}
174
175#[derive(Clone)]
176struct QueryRetryState {
177 resume_after: String,
178 remaining_limit: Option<usize>,
179 offset: usize,
180}
181
182async fn fetch_examples_with_query<MakeBindings>(
183 http_client: Arc<dyn HttpClient>,
184 step_progress: &crate::progress::StepProgress,
185 background_executor: BackgroundExecutor,
186 statement: &str,
187 initial_retry_state: QueryRetryState,
188 make_bindings: MakeBindings,
189 required_columns: &[&str],
190 parse_response: for<'a> fn(
191 &'a SnowflakeStatementResponse,
192 &'a HashMap<String, usize>,
193 ) -> Result<Box<dyn Iterator<Item = Example> + 'a>>,
194) -> Result<Vec<Example>>
195where
196 MakeBindings: Fn(&QueryRetryState) -> JsonValue,
197{
198 let snowflake = SnowflakeConfig {
199 token: std::env::var("EP_SNOWFLAKE_API_KEY")
200 .context("missing required environment variable EP_SNOWFLAKE_API_KEY")?,
201 base_url: std::env::var("EP_SNOWFLAKE_BASE_URL").context(
202 "missing required environment variable EP_SNOWFLAKE_BASE_URL (e.g. https://<account>.snowflakecomputing.com)",
203 )?,
204 role: std::env::var("EP_SNOWFLAKE_ROLE").ok(),
205 };
206
207 let mut requested_columns = required_columns.to_vec();
208 if !requested_columns.contains(&"continuation_time") {
209 requested_columns.push("continuation_time");
210 }
211
212 let mut parsed_examples = Vec::new();
213 let mut retry_state = initial_retry_state;
214 let mut retry_count = 0usize;
215
216 loop {
217 let bindings = make_bindings(&retry_state);
218 let request = json!({
219 "statement": statement,
220 "database": "EVENTS",
221 "schema": "PUBLIC",
222 "warehouse": "DBT",
223 "role": snowflake.role.as_deref(),
224 "bindings": bindings
225 });
226
227 let response = match run_sql_with_polling(
228 http_client.clone(),
229 &snowflake.base_url,
230 &snowflake.token,
231 &request,
232 step_progress,
233 background_executor.clone(),
234 )
235 .await
236 {
237 Ok(response) => response,
238 Err(error) => {
239 if is_snowflake_timeout_error(&error) && !parsed_examples.is_empty() {
240 retry_count += 1;
241 step_progress.set_substatus(format!(
242 "retrying from {} ({retry_count})",
243 retry_state.resume_after
244 ));
245 continue;
246 }
247
248 return Err(error);
249 }
250 };
251
252 let total_rows = response
253 .result_set_meta_data
254 .as_ref()
255 .and_then(|meta| meta.num_rows)
256 .unwrap_or(response.data.len() as i64);
257 let partition_count = response
258 .result_set_meta_data
259 .as_ref()
260 .map(|meta| meta.partition_info.len())
261 .unwrap_or(1)
262 .max(1);
263
264 step_progress.set_info(format!("{} rows", total_rows), InfoStyle::Normal);
265 step_progress.set_substatus("parsing");
266
267 let column_indices = get_column_indices(&response.result_set_meta_data, &requested_columns);
268 let mut rows_fetched_this_attempt = 0usize;
269 let mut timed_out_fetching_partition = false;
270
271 parsed_examples.extend(parse_response(&response, &column_indices)?);
272 rows_fetched_this_attempt += response.data.len();
273 let mut last_continuation_time_this_attempt =
274 last_continuation_timestamp_from_response(&response, &column_indices);
275
276 if partition_count > 1 {
277 let statement_handle = response
278 .statement_handle
279 .as_ref()
280 .context("response has multiple partitions but no statementHandle")?;
281
282 for partition in 1..partition_count {
283 step_progress.set_substatus(format!(
284 "fetching partition {}/{}",
285 partition + 1,
286 partition_count
287 ));
288
289 let partition_response = match fetch_partition_with_retries(
290 http_client.clone(),
291 &snowflake.base_url,
292 &snowflake.token,
293 statement_handle,
294 partition,
295 background_executor.clone(),
296 )
297 .await
298 {
299 Ok(response) => response,
300 Err(error) => {
301 if is_snowflake_timeout_error(&error) && rows_fetched_this_attempt > 0 {
302 timed_out_fetching_partition = true;
303 break;
304 }
305
306 return Err(error);
307 }
308 };
309
310 parsed_examples.extend(parse_response(&partition_response, &column_indices)?);
311 rows_fetched_this_attempt += partition_response.data.len();
312
313 if let Some(partition_continuation_time) =
314 last_continuation_timestamp_from_response(&partition_response, &column_indices)
315 {
316 last_continuation_time_this_attempt = Some(partition_continuation_time);
317 }
318 }
319 }
320
321 if rows_fetched_this_attempt == 0 {
322 step_progress.set_substatus("done");
323 return Ok(parsed_examples);
324 }
325
326 if let Some(remaining_limit_value) = &mut retry_state.remaining_limit {
327 *remaining_limit_value =
328 remaining_limit_value.saturating_sub(rows_fetched_this_attempt);
329 if *remaining_limit_value == 0 {
330 step_progress.set_substatus("done");
331 return Ok(parsed_examples);
332 }
333 }
334
335 if !timed_out_fetching_partition {
336 step_progress.set_substatus("done");
337 return Ok(parsed_examples);
338 }
339
340 let Some(last_continuation_time_this_attempt) = last_continuation_time_this_attempt else {
341 step_progress.set_substatus("done");
342 return Ok(parsed_examples);
343 };
344
345 retry_state.resume_after = last_continuation_time_this_attempt;
346 retry_state.offset = 0;
347 retry_count += 1;
348 step_progress.set_substatus(format!(
349 "retrying from {} ({retry_count})",
350 retry_state.resume_after
351 ));
352 }
353}
354
355pub(crate) async fn fetch_partition(
356 http_client: Arc<dyn HttpClient>,
357 base_url: &str,
358 token: &str,
359 statement_handle: &str,
360 partition: usize,
361) -> Result<SnowflakeStatementResponse> {
362 let url = format!(
363 "{}/api/v2/statements/{}?partition={}",
364 base_url.trim_end_matches('/'),
365 statement_handle,
366 partition
367 );
368
369 let http_request = Request::builder()
370 .method(Method::GET)
371 .uri(url.as_str())
372 .header("Authorization", format!("Bearer {token}"))
373 .header(
374 "X-Snowflake-Authorization-Token-Type",
375 "PROGRAMMATIC_ACCESS_TOKEN",
376 )
377 .header("Accept", "application/json")
378 .header("Accept-Encoding", "gzip")
379 .header("User-Agent", "edit_prediction_cli")
380 .body(AsyncBody::empty())?;
381
382 let response = http_client
383 .send(http_request)
384 .await
385 .context("failed to send partition request to Snowflake SQL API")?;
386
387 let status = response.status();
388 let content_encoding = response
389 .headers()
390 .get("content-encoding")
391 .and_then(|v| v.to_str().ok())
392 .map(|s| s.to_lowercase());
393
394 let body_bytes = {
395 use futures::AsyncReadExt as _;
396
397 let mut body = response.into_body();
398 let mut bytes = Vec::new();
399 body.read_to_end(&mut bytes)
400 .await
401 .context("failed to read Snowflake SQL API partition response body")?;
402 bytes
403 };
404
405 let body_bytes = if content_encoding.as_deref() == Some("gzip") {
406 let mut decoder = GzDecoder::new(&body_bytes[..]);
407 let mut decompressed = Vec::new();
408 decoder
409 .read_to_end(&mut decompressed)
410 .context("failed to decompress gzip response")?;
411 decompressed
412 } else {
413 body_bytes
414 };
415
416 if !status.is_success() && status.as_u16() != 202 {
417 let body_text = String::from_utf8_lossy(&body_bytes);
418 anyhow::bail!(
419 "snowflake sql api partition request http {}: {}",
420 status.as_u16(),
421 body_text
422 );
423 }
424
425 if body_bytes.is_empty() {
426 anyhow::bail!(
427 "snowflake sql api partition {} returned empty response body (http {})",
428 partition,
429 status.as_u16()
430 );
431 }
432
433 serde_json::from_slice::<SnowflakeStatementResponse>(&body_bytes).with_context(|| {
434 let body_preview = String::from_utf8_lossy(&body_bytes[..body_bytes.len().min(500)]);
435 format!(
436 "failed to parse Snowflake SQL API partition {} response JSON (http {}): {}",
437 partition,
438 status.as_u16(),
439 body_preview
440 )
441 })
442}
443
444async fn fetch_partition_with_retries(
445 http_client: Arc<dyn HttpClient>,
446 base_url: &str,
447 token: &str,
448 statement_handle: &str,
449 partition: usize,
450 background_executor: BackgroundExecutor,
451) -> Result<SnowflakeStatementResponse> {
452 let mut last_error = None;
453
454 for retry_attempt in 0..=PARTITION_FETCH_MAX_RETRIES {
455 match fetch_partition(
456 http_client.clone(),
457 base_url,
458 token,
459 statement_handle,
460 partition,
461 )
462 .await
463 {
464 Ok(response) => return Ok(response),
465 Err(error) => {
466 if retry_attempt == PARTITION_FETCH_MAX_RETRIES
467 || !is_transient_partition_fetch_error(&error)
468 {
469 return Err(error);
470 }
471
472 last_error = Some(error);
473 background_executor
474 .timer(PARTITION_FETCH_RETRY_DELAYS[retry_attempt])
475 .await;
476 }
477 }
478 }
479
480 match last_error {
481 Some(error) => Err(error),
482 None => anyhow::bail!("partition fetch retry loop exited without a result"),
483 }
484}
485
486fn is_transient_partition_fetch_error(error: &anyhow::Error) -> bool {
487 error.chain().any(|cause| {
488 let message = cause.to_string();
489 message.contains("failed to read Snowflake SQL API partition response body")
490 || message.contains("unexpected EOF")
491 || message.contains("peer closed connection without sending TLS close_notify")
492 })
493}
494
495pub(crate) async fn run_sql(
496 http_client: Arc<dyn HttpClient>,
497 base_url: &str,
498 token: &str,
499 request: &serde_json::Value,
500) -> Result<SnowflakeStatementResponse> {
501 let url = format!("{}/api/v2/statements", base_url.trim_end_matches('/'));
502
503 let request_body =
504 serde_json::to_vec(request).context("failed to serialize Snowflake SQL API request")?;
505
506 let http_request = Request::builder()
507 .method(Method::POST)
508 .uri(url.as_str())
509 .header("Authorization", format!("Bearer {token}"))
510 .header(
511 "X-Snowflake-Authorization-Token-Type",
512 "PROGRAMMATIC_ACCESS_TOKEN",
513 )
514 .header("Content-Type", "application/json")
515 .header("Accept", "application/json")
516 .header("User-Agent", "edit_prediction_cli")
517 .body(AsyncBody::from(request_body.clone()))?;
518
519 let response = http_client
520 .send(http_request)
521 .await
522 .context("failed to send request to Snowflake SQL API")?;
523
524 let status = response.status();
525 let body_bytes = {
526 use futures::AsyncReadExt as _;
527
528 let mut body = response.into_body();
529 let mut bytes = Vec::new();
530 body.read_to_end(&mut bytes)
531 .await
532 .context("failed to read Snowflake SQL API response body")?;
533 bytes
534 };
535
536 let snowflake_response = serde_json::from_slice::<SnowflakeStatementResponse>(&body_bytes)
537 .context("failed to parse Snowflake SQL API response JSON")?;
538
539 if !status.is_success() && status.as_u16() != 202 && !is_timeout_response(&snowflake_response) {
540 let body_text = String::from_utf8_lossy(&body_bytes);
541 anyhow::bail!("snowflake sql api http {}: {}", status.as_u16(), body_text);
542 }
543
544 if is_timeout_response(&snowflake_response) {
545 anyhow::bail!(
546 "snowflake sql api timed out code={} message={}",
547 snowflake_response.code.as_deref().unwrap_or("<no code>"),
548 snowflake_response
549 .message
550 .as_deref()
551 .unwrap_or("<no message>")
552 );
553 }
554
555 Ok(snowflake_response)
556}
557
558pub async fn fetch_rejected_examples_after(
559 http_client: Arc<dyn HttpClient>,
560 after_timestamps: &[(bool, String)],
561 max_rows_per_timestamp: Option<usize>,
562 offset: usize,
563 background_executor: BackgroundExecutor,
564 min_capture_version: Option<MinCaptureVersion>,
565) -> Result<Vec<Example>> {
566 if after_timestamps.is_empty() {
567 return Ok(Vec::new());
568 }
569
570 let progress = Progress::global();
571
572 let mut all_examples = Vec::new();
573
574 for (explicit, after_date) in after_timestamps.iter() {
575 let step_progress_name = format!("rejected>{after_date}");
576 let step_progress = progress.start(Step::PullExamples, &step_progress_name);
577 step_progress.set_substatus("querying");
578
579 let min_version_str = min_capture_version.map(|version| {
580 (version.major as u64 * 1_000_000 + version.minor as u64 * 1_000 + version.patch as u64)
581 .to_string()
582 });
583 let min_version_ref = min_version_str.as_deref();
584
585 let statement = indoc! {r#"
586 SELECT
587 ep_request_id AS request_id,
588 device_id AS device_id,
589 requested_at::string AS continuation_time,
590 requested_at::string AS time,
591 input_payload AS input,
592 prompt AS prompt,
593 requested_output AS output,
594 settled_editable_region AS settled_editable_region,
595 is_ep_shown_before_rejected AS was_shown,
596 ep_rejected_reason AS reason,
597 zed_version AS zed_version
598 FROM ZED_DBT.DBT_PROD.fct_edit_prediction_examples
599 WHERE ep_outcome LIKE ?
600 AND is_ep_shown_before_rejected = true
601 AND requested_at > TRY_TO_TIMESTAMP_NTZ(?)
602 AND (? IS NULL OR (
603 COALESCE(TRY_CAST(SPLIT_PART(zed_version, '.', 1) AS INTEGER), 0) * 1000000
604 + COALESCE(TRY_CAST(SPLIT_PART(zed_version, '.', 2) AS INTEGER), 0) * 1000
605 + COALESCE(TRY_CAST(SPLIT_PART(SPLIT_PART(zed_version, '.', 3), '+', 1) AS INTEGER), 0)
606 ) >= ?)
607 ORDER BY requested_at ASC
608 LIMIT ?
609 OFFSET ?
610 "#};
611
612 let examples = fetch_examples_with_query(
613 http_client.clone(),
614 &step_progress,
615 background_executor.clone(),
616 statement,
617 QueryRetryState {
618 resume_after: after_date.clone(),
619 remaining_limit: max_rows_per_timestamp,
620 offset,
621 },
622 |retry_state| {
623 json!({
624 "1": { "type": "TEXT", "value": if *explicit { "Rejected (Explicit)" } else { "Rejected%" } },
625 "2": { "type": "TEXT", "value": retry_state.resume_after },
626 "3": { "type": "FIXED", "value": min_version_ref },
627 "4": { "type": "FIXED", "value": min_version_ref },
628 "5": { "type": "FIXED", "value": format_limit(retry_state.remaining_limit) },
629 "6": { "type": "FIXED", "value": retry_state.offset.to_string() }
630 })
631 },
632 &[
633 "request_id",
634 "device_id",
635 "time",
636 "input",
637 "prompt",
638 "output",
639 "settled_editable_region",
640 "was_shown",
641 "reason",
642 "zed_version",
643 ],
644 rejected_examples_from_response,
645 )
646 .await?;
647
648 all_examples.extend(examples);
649 }
650
651 Ok(all_examples)
652}
653
654pub async fn fetch_accepted_examples_after(
655 http_client: Arc<dyn HttpClient>,
656 after_timestamps: &[String],
657 max_rows_per_timestamp: Option<usize>,
658 offset: usize,
659 background_executor: BackgroundExecutor,
660 min_capture_version: Option<MinCaptureVersion>,
661) -> Result<Vec<Example>> {
662 if after_timestamps.is_empty() {
663 return Ok(Vec::new());
664 }
665
666 let progress = Progress::global();
667
668 let mut all_examples = Vec::new();
669
670 for after_date in after_timestamps.iter() {
671 let step_progress_name = format!("accepted>{after_date}");
672 let step_progress = progress.start(Step::PullExamples, &step_progress_name);
673 step_progress.set_substatus("querying");
674
675 let min_version_str = min_capture_version.map(|version| {
676 (version.major as u64 * 1_000_000 + version.minor as u64 * 1_000 + version.patch as u64)
677 .to_string()
678 });
679 let min_version_ref = min_version_str.as_deref();
680
681 let statement = indoc! {r#"
682 SELECT
683 ep_request_id AS request_id,
684 device_id AS device_id,
685 requested_at::string AS continuation_time,
686 requested_at::string AS time,
687 input_payload AS input,
688 prompt AS prompt,
689 requested_output AS output,
690 settled_editable_region AS settled_editable_region,
691 zed_version AS zed_version
692 FROM ZED_DBT.DBT_PROD.fct_edit_prediction_examples
693 WHERE ep_outcome = 'Accepted'
694 AND requested_at > TRY_TO_TIMESTAMP_NTZ(?)
695 AND (? IS NULL OR (
696 COALESCE(TRY_CAST(SPLIT_PART(zed_version, '.', 1) AS INTEGER), 0) * 1000000
697 + COALESCE(TRY_CAST(SPLIT_PART(zed_version, '.', 2) AS INTEGER), 0) * 1000
698 + COALESCE(TRY_CAST(SPLIT_PART(SPLIT_PART(zed_version, '.', 3), '+', 1) AS INTEGER), 0)
699 ) >= ?)
700 ORDER BY requested_at ASC
701 LIMIT ?
702 OFFSET ?
703 "#};
704
705 let examples = fetch_examples_with_query(
706 http_client.clone(),
707 &step_progress,
708 background_executor.clone(),
709 statement,
710 QueryRetryState {
711 resume_after: after_date.clone(),
712 remaining_limit: max_rows_per_timestamp,
713 offset,
714 },
715 |retry_state| {
716 json!({
717 "1": { "type": "TEXT", "value": retry_state.resume_after },
718 "2": { "type": "FIXED", "value": min_version_ref },
719 "3": { "type": "FIXED", "value": min_version_ref },
720 "4": { "type": "FIXED", "value": format_limit(retry_state.remaining_limit) },
721 "5": { "type": "FIXED", "value": retry_state.offset.to_string() }
722 })
723 },
724 &[
725 "request_id",
726 "device_id",
727 "time",
728 "input",
729 "prompt",
730 "output",
731 "settled_editable_region",
732 "zed_version",
733 ],
734 accepted_examples_from_response,
735 )
736 .await?;
737
738 all_examples.extend(examples);
739 }
740
741 Ok(all_examples)
742}
743
744fn format_limit(limit: Option<usize>) -> String {
745 return limit.map(|l| l.to_string()).unwrap_or("NULL".to_string());
746}
747
748pub async fn fetch_requested_examples_after(
749 http_client: Arc<dyn HttpClient>,
750 after_timestamps: &[String],
751 max_rows_per_timestamp: Option<usize>,
752 offset: usize,
753 background_executor: BackgroundExecutor,
754 min_capture_version: Option<MinCaptureVersion>,
755) -> Result<Vec<Example>> {
756 if after_timestamps.is_empty() {
757 return Ok(Vec::new());
758 }
759
760 let progress = Progress::global();
761
762 let mut all_examples = Vec::new();
763
764 for after_date in after_timestamps.iter() {
765 let step_progress_name = format!("requested>{after_date}");
766 let step_progress = progress.start(Step::PullExamples, &step_progress_name);
767 step_progress.set_substatus("querying");
768
769 let min_version_str = min_capture_version.map(|version| {
770 (version.major as u64 * 1_000_000 + version.minor as u64 * 1_000 + version.patch as u64)
771 .to_string()
772 });
773 let min_version_ref = min_version_str.as_deref();
774
775 let statement = indoc! {r#"
776 SELECT
777 ep_request_id AS request_id,
778 device_id AS device_id,
779 requested_at::string AS continuation_time,
780 requested_at::string AS time,
781 input_payload AS input,
782 zed_version AS zed_version
783 FROM ZED_DBT.DBT_PROD.fct_edit_prediction_examples
784 WHERE requested_at > TRY_TO_TIMESTAMP_NTZ(?)
785 AND (? IS NULL OR (
786 COALESCE(TRY_CAST(SPLIT_PART(zed_version, '.', 1) AS INTEGER), 0) * 1000000
787 + COALESCE(TRY_CAST(SPLIT_PART(zed_version, '.', 2) AS INTEGER), 0) * 1000
788 + COALESCE(TRY_CAST(SPLIT_PART(SPLIT_PART(zed_version, '.', 3), '+', 1) AS INTEGER), 0)
789 ) >= ?)
790 ORDER BY requested_at ASC
791 LIMIT ?
792 OFFSET ?
793 "#};
794
795 let examples = fetch_examples_with_query(
796 http_client.clone(),
797 &step_progress,
798 background_executor.clone(),
799 statement,
800 QueryRetryState {
801 resume_after: after_date.clone(),
802 remaining_limit: max_rows_per_timestamp,
803 offset,
804 },
805 |retry_state| {
806 json!({
807 "1": { "type": "TEXT", "value": retry_state.resume_after },
808 "2": { "type": "FIXED", "value": min_version_ref },
809 "3": { "type": "FIXED", "value": min_version_ref },
810 "4": { "type": "FIXED", "value": format_limit(retry_state.remaining_limit) },
811 "5": { "type": "FIXED", "value": retry_state.offset.to_string() }
812 })
813 },
814 &["request_id", "device_id", "time", "input", "zed_version"],
815 requested_examples_from_response,
816 )
817 .await?;
818
819 all_examples.extend(examples);
820 }
821
822 Ok(all_examples)
823}
824
825pub async fn fetch_captured_examples_after(
826 http_client: Arc<dyn HttpClient>,
827 after_timestamps: &[String],
828 max_rows_per_timestamp: Option<usize>,
829 offset: usize,
830 background_executor: BackgroundExecutor,
831 min_capture_version: Option<MinCaptureVersion>,
832) -> Result<Vec<Example>> {
833 if after_timestamps.is_empty() {
834 return Ok(Vec::new());
835 }
836
837 let progress = Progress::global();
838
839 let mut all_examples = Vec::new();
840
841 for after_date in after_timestamps.iter() {
842 let step_progress_name = format!("captured>{after_date}");
843 let step_progress = progress.start(Step::PullExamples, &step_progress_name);
844 step_progress.set_substatus("querying");
845
846 let min_version_str = min_capture_version.map(|version| {
847 (version.major as u64 * 1_000_000 + version.minor as u64 * 1_000 + version.patch as u64)
848 .to_string()
849 });
850 let min_version_ref = min_version_str.as_deref();
851
852 let statement = indoc! {r#"
853 SELECT
854 ep_request_id AS request_id,
855 device_id AS device_id,
856 requested_at::string AS continuation_time,
857 requested_at::string AS time,
858 input_payload AS input,
859 settled_editable_region AS settled_editable_region,
860 example_payload AS example,
861 zed_version AS zed_version
862 FROM ZED_DBT.DBT_PROD.fct_edit_prediction_examples
863 WHERE settled_editable_region IS NOT NULL
864 AND example_payload IS NOT NULL
865 AND requested_at > TRY_TO_TIMESTAMP_NTZ(?)
866 AND (? IS NULL OR (
867 COALESCE(TRY_CAST(SPLIT_PART(zed_version, '.', 1) AS INTEGER), 0) * 1000000
868 + COALESCE(TRY_CAST(SPLIT_PART(zed_version, '.', 2) AS INTEGER), 0) * 1000
869 + COALESCE(TRY_CAST(SPLIT_PART(SPLIT_PART(zed_version, '.', 3), '+', 1) AS INTEGER), 0)
870 ) >= ?)
871 ORDER BY requested_at ASC
872 LIMIT ?
873 OFFSET ?
874 "#};
875
876 let examples = fetch_examples_with_query(
877 http_client.clone(),
878 &step_progress,
879 background_executor.clone(),
880 statement,
881 QueryRetryState {
882 resume_after: after_date.clone(),
883 remaining_limit: max_rows_per_timestamp,
884 offset,
885 },
886 |retry_state| {
887 json!({
888 "1": { "type": "TEXT", "value": retry_state.resume_after },
889 "2": { "type": "FIXED", "value": min_version_ref },
890 "3": { "type": "FIXED", "value": min_version_ref },
891 "4": { "type": "FIXED", "value": format_limit(retry_state.remaining_limit) },
892 "5": { "type": "FIXED", "value": retry_state.offset.to_string() }
893 })
894 },
895 &[
896 "request_id",
897 "device_id",
898 "time",
899 "input",
900 "settled_editable_region",
901 "example",
902 "zed_version",
903 ],
904 captured_examples_from_response,
905 )
906 .await?;
907
908 all_examples.extend(examples);
909 }
910
911 Ok(all_examples)
912}
913
914pub async fn fetch_settled_examples_after(
915 http_client: Arc<dyn HttpClient>,
916 after_timestamps: &[String],
917 max_rows_per_timestamp: Option<usize>,
918 offset: usize,
919 background_executor: BackgroundExecutor,
920 min_capture_version: Option<MinCaptureVersion>,
921) -> Result<Vec<Example>> {
922 if after_timestamps.is_empty() {
923 return Ok(Vec::new());
924 }
925
926 let progress = Progress::global();
927
928 let mut all_examples = Vec::new();
929
930 for after_date in after_timestamps.iter() {
931 let step_progress_name = format!("settled>{after_date}");
932 let step_progress = progress.start(Step::PullExamples, &step_progress_name);
933 step_progress.set_substatus("querying");
934
935 let _ = min_capture_version;
936
937 let statement = indoc! {r#"
938 SELECT
939 ep_request_id AS request_id,
940 device_id AS device_id,
941 requested_at::string AS continuation_time,
942 requested_at::string AS time,
943 input_payload AS input,
944 requested_output AS requested_output,
945 settled_editable_region AS settled_editable_region,
946 requested_format AS requested_format,
947 zed_version AS zed_version
948 FROM ZED_DBT.DBT_PROD.fct_edit_prediction_examples
949 WHERE settled_editable_region IS NOT NULL
950 AND requested_at > TRY_TO_TIMESTAMP_NTZ(?)
951 ORDER BY requested_at ASC
952 LIMIT ?
953 OFFSET ?
954 "#};
955
956 let examples = fetch_examples_with_query(
957 http_client.clone(),
958 &step_progress,
959 background_executor.clone(),
960 statement,
961 QueryRetryState {
962 resume_after: after_date.clone(),
963 remaining_limit: max_rows_per_timestamp,
964 offset,
965 },
966 |retry_state| {
967 json!({
968 "1": { "type": "TEXT", "value": retry_state.resume_after },
969 "2": { "type": "FIXED", "value": format_limit(retry_state.remaining_limit) },
970 "3": { "type": "FIXED", "value": retry_state.offset.to_string() }
971 })
972 },
973 &[
974 "request_id",
975 "device_id",
976 "time",
977 "input",
978 "requested_output",
979 "settled_editable_region",
980 "requested_format",
981 "zed_version",
982 ],
983 settled_examples_from_response,
984 )
985 .await?;
986
987 all_examples.extend(examples);
988 }
989
990 Ok(all_examples)
991}
992
993pub async fn fetch_rated_examples_after(
994 http_client: Arc<dyn HttpClient>,
995 inputs: &[(String, Option<EditPredictionRating>)],
996 max_rows_per_timestamp: Option<usize>,
997 offset: usize,
998 background_executor: BackgroundExecutor,
999 _min_capture_version: Option<MinCaptureVersion>,
1000) -> Result<Vec<Example>> {
1001 if inputs.is_empty() {
1002 return Ok(Vec::new());
1003 }
1004
1005 let progress = Progress::global();
1006
1007 let mut all_examples = Vec::new();
1008
1009 for (after_date, rating_filter) in inputs.iter() {
1010 let filter_label = match rating_filter {
1011 None => "",
1012 Some(EditPredictionRating::Positive) => ":positive",
1013 Some(EditPredictionRating::Negative) => ":negative",
1014 };
1015 let step_progress_name = format!("rated{filter_label}>{after_date}");
1016 let step_progress = progress.start(Step::PullExamples, &step_progress_name);
1017 step_progress.set_substatus("querying");
1018
1019 let rating_value = rating_filter.as_ref().map(|rating| match rating {
1020 EditPredictionRating::Positive => "Positive",
1021 EditPredictionRating::Negative => "Negative",
1022 });
1023
1024 let statement = indoc! {r#"
1025 SELECT
1026 ep_request_id AS request_id,
1027 rated_inputs AS inputs,
1028 rated_output AS output,
1029 settled_editable_region AS settled_editable_region,
1030 rating AS rating,
1031 feedback AS feedback,
1032 device_id AS device_id,
1033 requested_at::string AS continuation_time,
1034 requested_at::string AS time,
1035 NULL AS experiment_name,
1036 NULL AS environment,
1037 zed_version AS zed_version
1038 FROM ZED_DBT.DBT_PROD.fct_edit_prediction_examples
1039 WHERE rating IS NOT NULL
1040 AND (? IS NULL OR rating = ?)
1041 AND requested_at > TRY_TO_TIMESTAMP_NTZ(?)
1042 AND rated_inputs IS NOT NULL
1043 AND rated_inputs:cursor_excerpt IS NOT NULL
1044 AND rated_output IS NOT NULL
1045 ORDER BY requested_at ASC
1046 LIMIT ?
1047 OFFSET ?
1048 "#};
1049
1050 let examples = fetch_examples_with_query(
1051 http_client.clone(),
1052 &step_progress,
1053 background_executor.clone(),
1054 statement,
1055 QueryRetryState {
1056 resume_after: after_date.clone(),
1057 remaining_limit: max_rows_per_timestamp,
1058 offset,
1059 },
1060 |retry_state| {
1061 json!({
1062 "1": { "type": "TEXT", "value": rating_value },
1063 "2": { "type": "TEXT", "value": rating_value },
1064 "3": { "type": "TEXT", "value": retry_state.resume_after },
1065 "4": { "type": "FIXED", "value": format_limit(retry_state.remaining_limit) },
1066 "5": { "type": "FIXED", "value": retry_state.offset.to_string() }
1067 })
1068 },
1069 &[
1070 "request_id",
1071 "inputs",
1072 "output",
1073 "settled_editable_region",
1074 "rating",
1075 "feedback",
1076 "device_id",
1077 "time",
1078 "experiment_name",
1079 "environment",
1080 "zed_version",
1081 ],
1082 rated_examples_from_response,
1083 )
1084 .await?;
1085
1086 all_examples.extend(examples);
1087 }
1088
1089 Ok(all_examples)
1090}
1091
1092fn rated_examples_from_response<'a>(
1093 response: &'a SnowflakeStatementResponse,
1094 column_indices: &'a std::collections::HashMap<String, usize>,
1095) -> Result<Box<dyn Iterator<Item = Example> + 'a>> {
1096 if let Some(code) = &response.code {
1097 if code != SNOWFLAKE_SUCCESS_CODE {
1098 anyhow::bail!(
1099 "snowflake sql api returned error code={code} message={}",
1100 response.message.as_deref().unwrap_or("<no message>")
1101 );
1102 }
1103 }
1104
1105 let iter = response
1106 .data
1107 .iter()
1108 .enumerate()
1109 .filter_map(move |(row_index, data_row)| {
1110 let get_string = |name: &str| -> Option<String> {
1111 let index = column_indices.get(name).copied()?;
1112 match data_row.get(index)? {
1113 JsonValue::String(s) => Some(s.clone()),
1114 JsonValue::Null => None,
1115 other => Some(other.to_string()),
1116 }
1117 };
1118
1119 let get_json = |name: &str| -> Option<JsonValue> {
1120 let index = column_indices.get(name).copied()?;
1121 let value = data_row.get(index)?;
1122 if value.is_null() {
1123 return None;
1124 }
1125 match value {
1126 JsonValue::String(s) => serde_json::from_str(s).ok(),
1127 other => Some(other.clone()),
1128 }
1129 };
1130
1131 let request_id = get_string("request_id");
1132 let inputs_json = get_json("inputs");
1133 let inputs: Option<Zeta2PromptInput> = match &inputs_json {
1134 Some(v) => match serde_json::from_value(v.clone()) {
1135 Ok(parsed) => Some(parsed),
1136 Err(e) => {
1137 log::warn!(
1138 "skipping row {row_index}: failed to parse inputs - {e}",
1139 );
1140 return None;
1141 }
1142 },
1143 None => None,
1144 };
1145 let output = get_string("output");
1146 let settled_editable_region = get_string("settled_editable_region");
1147 let rating = get_string("rating");
1148 let feedback = get_string("feedback").unwrap_or_default();
1149 let device_id = get_string("device_id");
1150 let time = get_string("time");
1151 let experiment_name = get_string("experiment_name");
1152 let environment = get_string("environment");
1153 let zed_version = get_string("zed_version");
1154
1155 match (inputs, output.clone(), rating.clone(), time.clone()) {
1156 (Some(inputs), Some(output), Some(rating), Some(time)) => {
1157 Some(build_rated_example(
1158 request_id,
1159 device_id.unwrap_or_default(),
1160 time,
1161 inputs,
1162 output,
1163 settled_editable_region,
1164 rating,
1165 feedback,
1166 experiment_name,
1167 environment,
1168 zed_version,
1169 ))
1170 }
1171 _ => {
1172 log::warn!(
1173 "skipping row {row_index}: missing fields - inputs={:?} output={:?} rating={:?} time={:?}",
1174 inputs_json.is_some(),
1175 output.is_some(),
1176 rating.is_some(),
1177 time.is_some(),
1178 );
1179 None
1180 }
1181 }
1182 });
1183
1184 Ok(Box::new(iter))
1185}
1186
1187fn build_rated_example(
1188 request_id: Option<String>,
1189 device_id: String,
1190 time: String,
1191 input: Zeta2PromptInput,
1192 output: String,
1193 settled_editable_region: Option<String>,
1194 rating: String,
1195 feedback: String,
1196 experiment_name: Option<String>,
1197 environment: Option<String>,
1198 zed_version: Option<String>,
1199) -> Example {
1200 let parsed_rating = if rating == "Positive" {
1201 EditPredictionRating::Positive
1202 } else {
1203 EditPredictionRating::Negative
1204 };
1205 let is_positive = parsed_rating == EditPredictionRating::Positive;
1206 let request_id = request_id.unwrap_or_else(|| format!("rated-{}-{}", device_id, time));
1207
1208 let mut tags = Vec::with_capacity(3);
1209 tags.push(if is_positive {
1210 "rated:positive".to_string()
1211 } else {
1212 "rated:negative".to_string()
1213 });
1214 if let Some(experiment) = experiment_name {
1215 tags.push(format!("experiment:{experiment}"));
1216 }
1217 if let Some(env) = environment {
1218 tags.push(format!("environment:{env}"));
1219 }
1220
1221 let expected_patch = settled_editable_region
1222 .as_ref()
1223 .map(|settled_editable_region| {
1224 build_output_patch(
1225 &input.cursor_path,
1226 input.cursor_excerpt.as_ref(),
1227 &input.excerpt_ranges.editable_350,
1228 settled_editable_region,
1229 )
1230 });
1231 let mut example =
1232 build_example_from_snowflake(request_id, device_id, time, input, tags, None, zed_version);
1233
1234 example.spec.rating = Some(parsed_rating);
1235
1236 if !feedback.is_empty() {
1237 example
1238 .spec
1239 .human_feedback
1240 .push(edit_prediction::example_spec::HumanFeedback { message: feedback });
1241 }
1242
1243 if let Some(expected_patch) = expected_patch {
1244 example.spec.expected_patches = vec![expected_patch];
1245 } else if is_positive {
1246 example.spec.expected_patches = vec![output.clone()];
1247 }
1248
1249 if !is_positive {
1250 example.spec.rejected_patch = Some(output);
1251 }
1252
1253 example
1254}
1255
1256fn requested_examples_from_response<'a>(
1257 response: &'a SnowflakeStatementResponse,
1258 column_indices: &'a std::collections::HashMap<String, usize>,
1259) -> Result<Box<dyn Iterator<Item = Example> + 'a>> {
1260 if let Some(code) = &response.code {
1261 if code != SNOWFLAKE_SUCCESS_CODE {
1262 anyhow::bail!(
1263 "snowflake sql api returned error code={code} message={}",
1264 response.message.as_deref().unwrap_or("<no message>")
1265 );
1266 }
1267 }
1268
1269 let iter = response
1270 .data
1271 .iter()
1272 .enumerate()
1273 .filter_map(move |(row_index, data_row)| {
1274 let get_string = |name: &str| -> Option<String> {
1275 let index = column_indices.get(name).copied()?;
1276 match data_row.get(index)? {
1277 JsonValue::String(s) => Some(s.clone()),
1278 JsonValue::Null => None,
1279 other => Some(other.to_string()),
1280 }
1281 };
1282
1283 let get_json = |name: &str| -> Option<JsonValue> {
1284 let index = column_indices.get(name).copied()?;
1285 let value = data_row.get(index)?;
1286 if value.is_null() {
1287 return None;
1288 }
1289 match value {
1290 JsonValue::String(s) => serde_json::from_str(s).ok(),
1291 other => Some(other.clone()),
1292 }
1293 };
1294
1295 let request_id_str = get_string("request_id");
1296 let device_id = get_string("device_id");
1297 let time = get_string("time");
1298 let input_json = get_json("input");
1299 let input: Option<Zeta2PromptInput> =
1300 input_json.clone().and_then(|v| serde_json::from_value(v).ok());
1301 let zed_version = get_string("zed_version");
1302
1303 match (request_id_str.clone(), device_id.clone(), time.clone(), input) {
1304 (Some(request_id), Some(device_id), Some(time), Some(input)) => {
1305 Some(build_example_from_snowflake(
1306 request_id,
1307 device_id,
1308 time,
1309 input,
1310 vec!["requested".to_string()],
1311 None,
1312 zed_version,
1313 ))
1314 }
1315 _ => {
1316 log::warn!(
1317 "skipping row {row_index}: missing fields - request_id={:?} device_id={:?} time={:?} input={:?}",
1318 request_id_str.is_some(),
1319 device_id.is_some(),
1320 time.is_some(),
1321 input_json.is_some(),
1322 );
1323 None
1324 }
1325 }
1326 });
1327
1328 Ok(Box::new(iter))
1329}
1330
1331fn settled_examples_from_response<'a>(
1332 response: &'a SnowflakeStatementResponse,
1333 column_indices: &'a std::collections::HashMap<String, usize>,
1334) -> Result<Box<dyn Iterator<Item = Example> + 'a>> {
1335 if let Some(code) = &response.code {
1336 if code != SNOWFLAKE_SUCCESS_CODE {
1337 anyhow::bail!(
1338 "snowflake sql api returned error code={code} message={}",
1339 response.message.as_deref().unwrap_or("<no message>")
1340 );
1341 }
1342 }
1343
1344 let iter = response
1345 .data
1346 .iter()
1347 .enumerate()
1348 .filter_map(move |(row_index, data_row)| {
1349 let get_value = |name: &str| -> Option<JsonValue> {
1350 let index = column_indices.get(name).copied()?;
1351 let value = data_row.get(index)?;
1352 if value.is_null() {
1353 None
1354 } else {
1355 Some(value.clone())
1356 }
1357 };
1358
1359 let get_string = |name: &str| -> Option<String> {
1360 match get_value(name)? {
1361 JsonValue::String(s) => Some(s),
1362 other => Some(other.to_string()),
1363 }
1364 };
1365
1366 let parse_json_value = |raw: Option<&JsonValue>| -> Option<JsonValue> {
1367 let value = raw?;
1368 match value {
1369 JsonValue::String(s) => serde_json::from_str::<JsonValue>(s).ok(),
1370 other => Some(other.clone()),
1371 }
1372 };
1373
1374 let request_id_str = get_string("request_id");
1375 let device_id = get_string("device_id");
1376 let time = get_string("time");
1377 let input_raw = get_value("input");
1378 let input_json = parse_json_value(input_raw.as_ref());
1379 let input: Option<Zeta2PromptInput> = input_json
1380 .as_ref()
1381 .and_then(|parsed| serde_json::from_value(parsed.clone()).ok());
1382 let requested_output = get_string("requested_output");
1383 let settled_editable_region = get_string("settled_editable_region");
1384 let requested_format =
1385 get_string("requested_format").and_then(|s| ZetaFormat::parse(&s).ok());
1386 let zed_version = get_string("zed_version");
1387
1388 match (
1389 request_id_str.clone(),
1390 device_id.clone(),
1391 time.clone(),
1392 input.clone(),
1393 requested_output.clone(),
1394 settled_editable_region.clone(),
1395 requested_format,
1396 ) {
1397 (
1398 Some(request_id),
1399 Some(device_id),
1400 Some(time),
1401 Some(input),
1402 Some(requested_output),
1403 Some(settled_editable_region),
1404 Some(requested_format),
1405 ) => Some(build_settled_example(
1406 request_id,
1407 device_id,
1408 time,
1409 input,
1410 requested_output,
1411 settled_editable_region,
1412 requested_format,
1413 zed_version,
1414 )),
1415 _ => {
1416 let mut missing_fields = Vec::new();
1417
1418 if request_id_str.is_none() {
1419 missing_fields.push("request_id");
1420 }
1421 if device_id.is_none() {
1422 missing_fields.push("device_id");
1423 }
1424 if time.is_none() {
1425 missing_fields.push("time");
1426 }
1427 if input_raw.is_none() || input_json.is_none() || input.is_none() {
1428 missing_fields.push("input");
1429 }
1430 if requested_output.is_none() {
1431 missing_fields.push("requested_output");
1432 }
1433 if settled_editable_region.is_none() {
1434 missing_fields.push("settled_editable_region");
1435 }
1436 if requested_format.is_none() {
1437 missing_fields.push("requested_format");
1438 }
1439
1440 log::warn!(
1441 "skipping settled row {row_index}: [{}]",
1442 missing_fields.join(", "),
1443 );
1444 None
1445 }
1446 }
1447 });
1448
1449 Ok(Box::new(iter))
1450}
1451
1452fn captured_examples_from_response<'a>(
1453 response: &'a SnowflakeStatementResponse,
1454 column_indices: &'a std::collections::HashMap<String, usize>,
1455) -> Result<Box<dyn Iterator<Item = Example> + 'a>> {
1456 if let Some(code) = &response.code {
1457 if code != SNOWFLAKE_SUCCESS_CODE {
1458 anyhow::bail!(
1459 "snowflake sql api returned error code={code} message={}",
1460 response.message.as_deref().unwrap_or("<no message>")
1461 );
1462 }
1463 }
1464
1465 let iter = response
1466 .data
1467 .iter()
1468 .enumerate()
1469 .filter_map(move |(row_index, data_row)| {
1470 let get_value = |name: &str| -> Option<JsonValue> {
1471 let index = column_indices.get(name).copied()?;
1472 let value = data_row.get(index)?;
1473 if value.is_null() {
1474 None
1475 } else {
1476 Some(value.clone())
1477 }
1478 };
1479
1480 let get_string = |name: &str| -> Option<String> {
1481 match get_value(name)? {
1482 JsonValue::String(s) => Some(s),
1483 other => Some(other.to_string()),
1484 }
1485 };
1486
1487 let parse_json_value = |raw: Option<&JsonValue>| -> Option<JsonValue> {
1488 let value = raw?;
1489 match value {
1490 JsonValue::String(s) => serde_json::from_str::<JsonValue>(s).ok(),
1491 other => Some(other.clone()),
1492 }
1493 };
1494
1495 let request_id = get_string("request_id");
1496 let device_id = get_string("device_id");
1497 let time = get_string("time");
1498 let input_raw = get_value("input");
1499 let input_json = parse_json_value(input_raw.as_ref());
1500 let input: Option<Zeta2PromptInput> = input_json
1501 .as_ref()
1502 .and_then(|parsed| serde_json::from_value(parsed.clone()).ok());
1503 let example_raw = get_value("example");
1504 let example_json = parse_json_value(example_raw.as_ref());
1505 let example_spec: Option<ExampleSpec> = example_json.as_ref().and_then(|parsed| {
1506 serde_json::from_value(parsed.clone())
1507 .or_else(|_| {
1508 parsed
1509 .as_str()
1510 .and_then(|markdown| ExampleSpec::from_markdown(markdown).ok())
1511 .ok_or_else(|| {
1512 serde_json::Error::io(std::io::Error::other("not markdown"))
1513 })
1514 })
1515 .ok()
1516 });
1517 let has_example_spec = example_spec.is_some();
1518 let settled_editable_region = get_string("settled_editable_region");
1519 let zed_version = get_string("zed_version");
1520
1521 match (
1522 request_id.clone(),
1523 device_id.clone(),
1524 time.clone(),
1525 input.clone(),
1526 example_spec,
1527 settled_editable_region.clone(),
1528 ) {
1529 (
1530 Some(request_id),
1531 Some(device_id),
1532 Some(time),
1533 Some(input),
1534 Some(example_spec),
1535 Some(settled_editable_region),
1536 ) => Some(build_captured_example(
1537 request_id,
1538 device_id,
1539 time,
1540 input,
1541 example_spec,
1542 settled_editable_region,
1543 zed_version,
1544 )),
1545 _ => {
1546 let mut missing_fields = Vec::new();
1547
1548 if request_id.is_none() {
1549 missing_fields.push("request_id");
1550 }
1551 if device_id.is_none() {
1552 missing_fields.push("device_id");
1553 }
1554 if time.is_none() {
1555 missing_fields.push("time");
1556 }
1557 if input_raw.is_none() || input_json.is_none() || input.is_none() {
1558 missing_fields.push("input");
1559 }
1560 if example_raw.is_none() || !has_example_spec {
1561 missing_fields.push("example");
1562 }
1563 if settled_editable_region.is_none() {
1564 missing_fields.push("settled_editable_region");
1565 }
1566
1567 log::warn!(
1568 "skipping captured row {row_index}: [{}]",
1569 missing_fields.join(", "),
1570 );
1571 None
1572 }
1573 }
1574 });
1575
1576 Ok(Box::new(iter))
1577}
1578
1579fn build_settled_example(
1580 request_id: String,
1581 device_id: String,
1582 time: String,
1583 input: Zeta2PromptInput,
1584 requested_output: String,
1585 settled_editable_region: String,
1586 requested_format: ZetaFormat,
1587 zed_version: Option<String>,
1588) -> Example {
1589 let requested_editable_range =
1590 excerpt_range_for_format(requested_format, &input.excerpt_ranges).0;
1591
1592 let base_cursor_excerpt = input.cursor_excerpt.to_string();
1593
1594 let requested_range_is_valid = requested_editable_range.start <= requested_editable_range.end
1595 && requested_editable_range.end <= base_cursor_excerpt.len();
1596 let mut example = build_example_from_snowflake(
1597 request_id.clone(),
1598 device_id,
1599 time,
1600 input,
1601 vec!["settled".to_string()],
1602 None,
1603 zed_version,
1604 );
1605
1606 if !requested_range_is_valid {
1607 log::warn!(
1608 "skipping malformed requested range for request {}: requested={:?} (base_len={})",
1609 request_id,
1610 requested_editable_range,
1611 base_cursor_excerpt.len(),
1612 );
1613 return example;
1614 }
1615
1616 let settled_replacement = settled_editable_region.as_str();
1617 let rejected_patch = build_output_patch(
1618 &example.spec.cursor_path,
1619 &base_cursor_excerpt,
1620 &requested_editable_range,
1621 &requested_output,
1622 );
1623 let expected_patch = build_output_patch(
1624 &example.spec.cursor_path,
1625 &base_cursor_excerpt,
1626 &requested_editable_range,
1627 settled_replacement,
1628 );
1629
1630 example.spec.expected_patches = vec![expected_patch];
1631 example.spec.rejected_patch = Some(rejected_patch);
1632 example
1633}
1634
1635fn build_captured_example(
1636 request_id: String,
1637 device_id: String,
1638 time: String,
1639 input: Zeta2PromptInput,
1640 mut example_spec: ExampleSpec,
1641 settled_editable_region: String,
1642 zed_version: Option<String>,
1643) -> Example {
1644 let expected_patch = build_output_patch(
1645 &input.cursor_path,
1646 input.cursor_excerpt.as_ref(),
1647 &input.excerpt_ranges.editable_350,
1648 settled_editable_region.as_str(),
1649 );
1650
1651 example_spec.expected_patches = vec![expected_patch];
1652 example_spec.telemetry = Some(TelemetrySource {
1653 request_id,
1654 device_id,
1655 time,
1656 rejection_reason: String::new(),
1657 was_shown: false,
1658 });
1659
1660 Example {
1661 spec: example_spec,
1662 zed_version,
1663 prompt_inputs: Some(input),
1664 prompt: None,
1665 predictions: Vec::new(),
1666 score: Vec::new(),
1667 qa: Vec::new(),
1668 state: None,
1669 }
1670}
1671
1672fn rejected_examples_from_response<'a>(
1673 response: &'a SnowflakeStatementResponse,
1674 column_indices: &'a std::collections::HashMap<String, usize>,
1675) -> Result<Box<dyn Iterator<Item = Example> + 'a>> {
1676 if let Some(code) = &response.code {
1677 if code != SNOWFLAKE_SUCCESS_CODE {
1678 anyhow::bail!(
1679 "snowflake sql api returned error code={code} message={}",
1680 response.message.as_deref().unwrap_or("<no message>")
1681 );
1682 }
1683 }
1684
1685 let iter = response
1686 .data
1687 .iter()
1688 .enumerate()
1689 .filter_map(move |(row_index, data_row)| {
1690 let get_string = |name: &str| -> Option<String> {
1691 let index = column_indices.get(name).copied()?;
1692 match data_row.get(index)? {
1693 JsonValue::String(s) => Some(s.clone()),
1694 JsonValue::Null => None,
1695 other => Some(other.to_string()),
1696 }
1697 };
1698
1699 let get_json = |name: &str| -> Option<JsonValue> {
1700 let index = column_indices.get(name).copied()?;
1701 let value = data_row.get(index)?;
1702 if value.is_null() {
1703 return None;
1704 }
1705 match value {
1706 JsonValue::String(s) => serde_json::from_str(s).ok(),
1707 other => Some(other.clone()),
1708 }
1709 };
1710
1711 let get_bool = |name: &str| -> Option<bool> {
1712 let index = column_indices.get(name).copied()?;
1713 match data_row.get(index)? {
1714 JsonValue::Bool(b) => Some(*b),
1715 JsonValue::String(s) => s.parse().ok(),
1716 _ => None,
1717 }
1718 };
1719
1720 let request_id_str = get_string("request_id");
1721 let device_id = get_string("device_id");
1722 let time = get_string("time");
1723 let input_json = get_json("input");
1724 let input: Option<Zeta2PromptInput> =
1725 input_json.clone().and_then(|v| serde_json::from_value(v).ok());
1726 let prompt = get_string("prompt");
1727 let output = get_string("output");
1728 let settled_editable_region = get_string("settled_editable_region");
1729 let was_shown = get_bool("was_shown");
1730 let reason = get_string("reason");
1731 let zed_version = get_string("zed_version");
1732
1733 match (request_id_str.clone(), device_id.clone(), time.clone(), input, output.clone(), was_shown, reason.clone()) {
1734 (Some(request_id), Some(device_id), Some(time), Some(input), Some(output), Some(was_shown), Some(reason)) => {
1735 Some(build_rejected_example(
1736 request_id,
1737 device_id,
1738 time,
1739 input,
1740 prompt,
1741 output,
1742 settled_editable_region,
1743 was_shown,
1744 reason,
1745 zed_version,
1746 ))
1747 }
1748 _ => {
1749 log::warn!(
1750 "skipping row {row_index}: missing fields - request_id={:?} device_id={:?} time={:?} input={:?} output={:?} was_shown={:?} reason={:?}",
1751 request_id_str.is_some(),
1752 device_id.is_some(),
1753 time.is_some(),
1754 input_json.is_some(),
1755 output.is_some(),
1756 was_shown.is_some(),
1757 reason.is_some()
1758 );
1759 None
1760 }
1761 }
1762 });
1763
1764 Ok(Box::new(iter))
1765}
1766
1767fn build_rejected_example(
1768 request_id: String,
1769 device_id: String,
1770 time: String,
1771 input: Zeta2PromptInput,
1772 prompt: Option<String>,
1773 output: String,
1774 settled_editable_region: Option<String>,
1775 was_shown: bool,
1776 reason: String,
1777 zed_version: Option<String>,
1778) -> Example {
1779 let rejected_patch = build_output_patch(
1780 &input.cursor_path,
1781 input.cursor_excerpt.as_ref(),
1782 &input.excerpt_ranges.editable_350,
1783 &output,
1784 );
1785 let expected_patch = settled_editable_region
1786 .as_ref()
1787 .map(|settled_editable_region| {
1788 build_output_patch(
1789 &input.cursor_path,
1790 input.cursor_excerpt.as_ref(),
1791 &input.excerpt_ranges.editable_350,
1792 settled_editable_region,
1793 )
1794 });
1795 let mut example = build_example_from_snowflake(
1796 request_id,
1797 device_id,
1798 time,
1799 input,
1800 vec![format!("rejection:{}", reason.to_lowercase())],
1801 Some(RejectionInfo { reason, was_shown }),
1802 zed_version,
1803 );
1804 example.spec.rejected_patch = Some(rejected_patch.clone());
1805 if let Some(expected_patch) = expected_patch {
1806 example.spec.expected_patches = vec![expected_patch];
1807 }
1808 example.predictions.push(ExamplePrediction {
1809 provider: PredictionProvider::default(),
1810 actual_output: output.clone(),
1811 actual_patch: Some(rejected_patch),
1812 actual_cursor: None,
1813 error: None,
1814 cumulative_logprob: None,
1815 avg_logprob: None,
1816 });
1817 example.prompt = prompt.map(|prompt| ExamplePrompt {
1818 input: prompt,
1819 expected_output: None,
1820 rejected_output: Some(output),
1821 prefill: None,
1822 provider: PredictionProvider::default(),
1823 });
1824 example
1825}
1826
1827struct RejectionInfo {
1828 reason: String,
1829 was_shown: bool,
1830}
1831
1832fn accepted_examples_from_response<'a>(
1833 response: &'a SnowflakeStatementResponse,
1834 column_indices: &'a std::collections::HashMap<String, usize>,
1835) -> Result<Box<dyn Iterator<Item = Example> + 'a>> {
1836 if let Some(code) = &response.code {
1837 if code != SNOWFLAKE_SUCCESS_CODE {
1838 anyhow::bail!(
1839 "snowflake sql api returned error code={code} message={}",
1840 response.message.as_deref().unwrap_or("<no message>")
1841 );
1842 }
1843 }
1844
1845 let iter = response
1846 .data
1847 .iter()
1848 .enumerate()
1849 .filter_map(move |(row_index, data_row)| {
1850 let get_string = |name: &str| -> Option<String> {
1851 let index = column_indices.get(name).copied()?;
1852 match data_row.get(index)? {
1853 JsonValue::String(s) => Some(s.clone()),
1854 JsonValue::Null => None,
1855 other => Some(other.to_string()),
1856 }
1857 };
1858
1859 let get_json = |name: &str| -> Option<JsonValue> {
1860 let index = column_indices.get(name).copied()?;
1861 let value = data_row.get(index)?;
1862 if value.is_null() {
1863 return None;
1864 }
1865 match value {
1866 JsonValue::String(s) => serde_json::from_str(s).ok(),
1867 other => Some(other.clone()),
1868 }
1869 };
1870
1871 let request_id_str = get_string("request_id");
1872 let device_id = get_string("device_id");
1873 let time = get_string("time");
1874 let input_json = get_json("input");
1875 let input: Option<Zeta2PromptInput> =
1876 input_json.clone().and_then(|v| serde_json::from_value(v).ok());
1877 let prompt = get_string("prompt");
1878 let output = get_string("output");
1879 let settled_editable_region = get_string("settled_editable_region");
1880 let zed_version = get_string("zed_version");
1881
1882 match (request_id_str.clone(), device_id.clone(), time.clone(), input, output.clone()) {
1883 (Some(request_id), Some(device_id), Some(time), Some(input), Some(output)) => {
1884 Some(build_accepted_example(
1885 request_id,
1886 device_id,
1887 time,
1888 input,
1889 prompt,
1890 output,
1891 settled_editable_region,
1892 zed_version,
1893 ))
1894 }
1895 _ => {
1896 log::warn!(
1897 "skipping row {row_index}: missing fields - request_id={:?} device_id={:?} time={:?} input={:?} output={:?}",
1898 request_id_str.is_some(),
1899 device_id.is_some(),
1900 time.is_some(),
1901 input_json.is_some(),
1902 output.is_some(),
1903 );
1904 None
1905 }
1906 }
1907 });
1908
1909 Ok(Box::new(iter))
1910}
1911
1912fn build_accepted_example(
1913 request_id: String,
1914 device_id: String,
1915 time: String,
1916 input: Zeta2PromptInput,
1917 prompt: Option<String>,
1918 output: String,
1919 settled_editable_region: Option<String>,
1920 zed_version: Option<String>,
1921) -> Example {
1922 let accepted_patch = build_output_patch(
1923 &input.cursor_path,
1924 input.cursor_excerpt.as_ref(),
1925 &input.excerpt_ranges.editable_350,
1926 &output,
1927 );
1928 let expected_patch = settled_editable_region
1929 .as_ref()
1930 .map(|settled_editable_region| {
1931 build_output_patch(
1932 &input.cursor_path,
1933 input.cursor_excerpt.as_ref(),
1934 &input.excerpt_ranges.editable_350,
1935 settled_editable_region,
1936 )
1937 });
1938 let mut example = build_example_from_snowflake(
1939 request_id,
1940 device_id,
1941 time,
1942 input,
1943 vec!["accepted".to_string()],
1944 None,
1945 zed_version,
1946 );
1947 if let Some(expected_patch) = expected_patch {
1948 example.spec.expected_patches = vec![expected_patch];
1949 }
1950 example.predictions.push(ExamplePrediction {
1951 provider: PredictionProvider::default(),
1952 actual_output: output.clone(),
1953 actual_patch: Some(accepted_patch),
1954 actual_cursor: None, // todo: why no cursor?
1955 error: None,
1956 cumulative_logprob: None,
1957 avg_logprob: None,
1958 });
1959 example.prompt = prompt.map(|prompt| ExamplePrompt {
1960 input: prompt,
1961 expected_output: Some(output),
1962 rejected_output: None,
1963 prefill: None,
1964 provider: PredictionProvider::default(),
1965 });
1966 example
1967}
1968
1969fn build_example_from_snowflake(
1970 request_id: String,
1971 device_id: String,
1972 time: String,
1973 input: Zeta2PromptInput,
1974 tags: Vec<String>,
1975 rejection: Option<RejectionInfo>,
1976 zed_version: Option<String>,
1977) -> Example {
1978 let cursor_excerpt = input.cursor_excerpt.as_ref();
1979 let cursor_offset = input.cursor_offset_in_excerpt;
1980
1981 let mut edit_history = String::new();
1982 for event in &input.events {
1983 zeta_prompt::write_event(&mut edit_history, event);
1984 edit_history.push('\n');
1985 }
1986
1987 let (rejection_reason, was_shown) = match &rejection {
1988 Some(r) => (r.reason.clone(), r.was_shown),
1989 None => (String::new(), false),
1990 };
1991
1992 let spec = ExampleSpec {
1993 name: request_id.clone(),
1994 repository_url: String::new(),
1995 revision: String::new(),
1996 tags,
1997 reasoning: None,
1998 uncommitted_diff: String::new(),
1999 recently_opened_files: Vec::new(),
2000 recently_viewed_files: Vec::new(),
2001 uncommitted_diff_contains_edit_history: false,
2002 cursor_path: input.cursor_path.clone(),
2003 cursor_position: build_cursor_position(cursor_excerpt, cursor_offset),
2004 edit_history,
2005 expected_patches: Vec::new(),
2006 rejected_patch: None,
2007 telemetry: Some(TelemetrySource {
2008 request_id,
2009 device_id,
2010 time,
2011 rejection_reason,
2012 was_shown,
2013 }),
2014 human_feedback: Vec::new(),
2015 rating: None,
2016 };
2017
2018 Example {
2019 spec,
2020 zed_version,
2021 prompt_inputs: Some(input),
2022 prompt: None,
2023 predictions: Vec::new(),
2024 score: Vec::new(),
2025 qa: Vec::new(),
2026 state: None,
2027 }
2028}
2029
2030fn build_cursor_position(excerpt: &str, cursor_offset: usize) -> String {
2031 let before = &excerpt[..cursor_offset.min(excerpt.len())];
2032 let after = &excerpt[cursor_offset.min(excerpt.len())..];
2033 format!("{}[CURSOR_POSITION]{}", before, after)
2034}
2035
2036fn build_output_patch(
2037 cursor_path: &std::path::Path,
2038 cursor_excerpt: &str,
2039 editable_range: &std::ops::Range<usize>,
2040 model_output: &str,
2041) -> String {
2042 let old_text = &cursor_excerpt[editable_range.clone()];
2043
2044 let editable_start_row = cursor_excerpt[..editable_range.start]
2045 .chars()
2046 .filter(|&c| c == '\n')
2047 .count() as u32;
2048
2049 let diff_body = language::unified_diff_with_offsets(
2050 old_text,
2051 model_output,
2052 editable_start_row,
2053 editable_start_row,
2054 );
2055
2056 let mut patch = String::new();
2057 writeln!(&mut patch, "--- a/{}", cursor_path.display()).ok();
2058 writeln!(&mut patch, "+++ b/{}", cursor_path.display()).ok();
2059 patch.push_str(&diff_body);
2060 patch
2061}
2062
2063fn is_timeout_response(response: &SnowflakeStatementResponse) -> bool {
2064 response.code.as_deref() == Some(SNOWFLAKE_TIMEOUT_CODE)
2065 && response
2066 .message
2067 .as_deref()
2068 .map(|message| message.to_ascii_lowercase().contains("timeout"))
2069 .unwrap_or(false)
2070}
2071
2072fn is_snowflake_timeout_error(error: &anyhow::Error) -> bool {
2073 error
2074 .chain()
2075 .any(|cause| cause.to_string().contains(SNOWFLAKE_TIMEOUT_CODE))
2076}
2077
2078fn last_continuation_timestamp_from_response(
2079 response: &SnowflakeStatementResponse,
2080 column_indices: &HashMap<String, usize>,
2081) -> Option<String> {
2082 let continuation_time_index = column_indices.get("continuation_time").copied()?;
2083 response
2084 .data
2085 .iter()
2086 .rev()
2087 .find_map(|row| match row.get(continuation_time_index)? {
2088 JsonValue::String(value) => Some(value.clone()),
2089 JsonValue::Null => None,
2090 other => Some(other.to_string()),
2091 })
2092}
2093
2094pub(crate) fn get_column_indices(
2095 meta: &Option<SnowflakeResultSetMetaData>,
2096 names: &[&str],
2097) -> HashMap<String, usize> {
2098 let mut indices = HashMap::new();
2099 if let Some(meta) = meta {
2100 for (index, col) in meta.row_type.iter().enumerate() {
2101 for &name in names {
2102 if col.name.eq_ignore_ascii_case(name) {
2103 indices.insert(name.to_string(), index);
2104 }
2105 }
2106 }
2107 }
2108 indices
2109}
2110