Skip to repository content

tenant.openagents/omega

No repository description is available.

OpenAgents Git authority 2026-07-28T04:57:58.883Z Public web read
NIP-34 coordinate30617:7649603503856e5148d571eac2766b288a8ff1e9e35d380337a1d2b0015b4f92:omega
MaintainersHidden in public view
References2 branches · 1 tag
Read-only clonegit clone https://openagents.com/git/tenant.openagents/omega.git
Browse files

pull_examples.rs

2110 lines · 73.3 KB · rust
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
Served at tenant.openagents/omega Member data and write actions are omitted.