Skip to repository content

tenant.openagents/omega

No repository description is available.

OpenAgents Git authority 2026-07-28T02:54:35.498Z 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

predict.rs

822 lines · 28.9 KB · rust
1use crate::{
2    FormatPromptArgs, PredictArgs, PredictionProvider, TeacherBackend,
3    anthropic_client::AnthropicClient,
4    example::{Example, ExamplePrediction, ExamplePrompt},
5    format_prompt::{TeacherJumpsPrompt, TeacherPrompt, run_format_prompt},
6    headless::EpAppState,
7    load_project::run_load_project,
8    openai_client::OpenAiClient,
9    parse_output::parse_prediction_output,
10    paths::{LATEST_EXAMPLE_RUN_DIR, RUN_DIR},
11    progress::{ExampleProgress, InfoStyle, Progress, Step, StepProgress},
12    retrieve_context::{ContextRetrievalType, run_context_retrieval},
13};
14use anyhow::Context as _;
15use cloud_llm_client::predict_edits_v3::{RawCompletionRequest, RawCompletionResponse};
16use edit_prediction::{DebugEvent, EditPredictionStore, Zeta2RawConfig};
17use futures::{AsyncReadExt as _, FutureExt as _, StreamExt as _, future::Shared};
18use gpui::{AppContext as _, AsyncApp, Task};
19use http_client::{AsyncBody, HttpClient, Method};
20use reqwest_client::ReqwestClient;
21use std::{
22    fs,
23    sync::{
24        Arc, Mutex, OnceLock,
25        atomic::{AtomicUsize, Ordering::SeqCst},
26    },
27};
28use zeta_prompt::ZetaFormat;
29
30static ANTHROPIC_CLIENT: OnceLock<AnthropicClient> = OnceLock::new();
31static OPENAI_CLIENT: OnceLock<OpenAiClient> = OnceLock::new();
32
33pub async fn run_prediction(
34    example: &mut Example,
35    args: &PredictArgs,
36    app_state: Arc<EpAppState>,
37    example_progress: &ExampleProgress,
38    mut cx: AsyncApp,
39) -> anyhow::Result<()> {
40    let repetition_count = args.repetitions;
41
42    if let Some(existing_prediction) = example.predictions.first() {
43        let has_prediction = existing_prediction.actual_patch.is_some()
44            || !existing_prediction.actual_output.is_empty();
45        if has_prediction {
46            match args.provider {
47                None => return Ok(()),
48                Some(provider) if existing_prediction.provider == provider => return Ok(()),
49                Some(_) => example.predictions.clear(),
50            }
51        }
52    }
53
54    let Some(provider) = args.provider else {
55        anyhow::bail!(
56            "No existing predictions found. Use --provider to specify which model to use for prediction."
57        );
58    };
59
60    if let PredictionProvider::Teacher(backend, _)
61    | PredictionProvider::TeacherNonBatching(backend, _)
62    | PredictionProvider::TeacherJumps(backend)
63    | PredictionProvider::TeacherJumpsNonBatching(backend) = provider
64    {
65        run_context_retrieval(
66            example,
67            app_state.clone(),
68            example_progress,
69            vec![ContextRetrievalType::Lsp],
70            false,
71            cx.clone(),
72        )
73        .await?;
74        run_format_prompt(
75            example,
76            &FormatPromptArgs {
77                provider,
78                related_files_budget: TeacherJumpsPrompt::DEFAULT_RELATED_FILES_BUDGET,
79            },
80            app_state.clone(),
81            example_progress,
82            cx,
83        )
84        .await?;
85
86        let step_progress = example_progress.start(Step::Predict);
87        let batched = matches!(
88            provider,
89            PredictionProvider::Teacher(..) | PredictionProvider::TeacherJumps(..)
90        );
91        return predict_teacher(
92            example,
93            backend,
94            batched,
95            repetition_count,
96            args.cache_only,
97            &step_progress,
98        )
99        .await;
100    }
101
102    if let PredictionProvider::Baseten(format) = provider {
103        run_format_prompt(
104            example,
105            &FormatPromptArgs {
106                provider: PredictionProvider::Zeta2(format),
107                related_files_budget: TeacherJumpsPrompt::DEFAULT_RELATED_FILES_BUDGET,
108            },
109            app_state.clone(),
110            example_progress,
111            cx,
112        )
113        .await?;
114
115        let step_progress = example_progress.start(Step::Predict);
116        return predict_baseten(example, format, &step_progress).await;
117    }
118
119    run_load_project(example, app_state.clone(), example_progress, cx.clone()).await?;
120    run_context_retrieval(
121        example,
122        app_state.clone(),
123        example_progress,
124        vec![ContextRetrievalType::Lsp],
125        false,
126        cx.clone(),
127    )
128    .await?;
129
130    let step_progress = example_progress.start(Step::Predict);
131
132    if matches!(
133        provider,
134        PredictionProvider::Zeta1 | PredictionProvider::Zeta2(_)
135    ) {
136        step_progress.set_substatus("authenticating");
137        static AUTHENTICATED: OnceLock<Shared<Task<()>>> = OnceLock::new();
138        AUTHENTICATED
139            .get_or_init(|| {
140                let client = app_state.client.clone();
141                cx.spawn(async move |cx| {
142                    if let Err(e) = client.sign_in_with_optional_connect(true, cx).await {
143                        eprintln!("Authentication failed: {}", e);
144                    }
145                })
146                .shared()
147            })
148            .clone()
149            .await;
150    }
151
152    let ep_store = cx
153        .update(|cx| EditPredictionStore::try_global(cx))
154        .context("EditPredictionStore not initialized")?;
155
156    ep_store.update(&mut cx, |store, _cx| {
157        let model = match provider {
158            PredictionProvider::Zeta1 => edit_prediction::EditPredictionModel::Zeta,
159            PredictionProvider::Zeta2(_) => edit_prediction::EditPredictionModel::Zeta,
160            PredictionProvider::Mercury => edit_prediction::EditPredictionModel::Mercury,
161            PredictionProvider::Teacher(..)
162            | PredictionProvider::TeacherJumps(..)
163            | PredictionProvider::TeacherNonBatching(..)
164            | PredictionProvider::TeacherJumpsNonBatching(..)
165            | PredictionProvider::Repair
166            | PredictionProvider::Baseten(_) => {
167                unreachable!()
168            }
169        };
170        store.set_edit_prediction_model(model);
171
172        // If user specified a non-default Zeta2 version, configure raw endpoint.
173        // ZED_ZETA_MODEL env var is optional.
174        if let PredictionProvider::Zeta2(format) = provider {
175            if format != ZetaFormat::default() {
176                let model_id = std::env::var("ZED_ZETA_MODEL").ok();
177                let environment = std::env::var("ZED_ZETA_ENVIRONMENT").ok();
178                store.set_zeta2_raw_config(Zeta2RawConfig {
179                    model_id,
180                    environment,
181                    format,
182                });
183            }
184        }
185    });
186    step_progress.set_substatus("configuring model");
187    let state = example.state.as_ref().context("state must be set")?;
188    let run_dir = RUN_DIR.join(&example.spec.name);
189
190    let updated_example = Arc::new(Mutex::new(example.clone()));
191    let current_run_ix = Arc::new(AtomicUsize::new(0));
192
193    let mut debug_rx = ep_store.update(&mut cx, |store, cx| store.debug_info(&state.project, cx));
194    let debug_task = cx.background_spawn({
195        let updated_example = updated_example.clone();
196        let current_run_ix = current_run_ix.clone();
197        let run_dir = run_dir.clone();
198        async move {
199            while let Some(event) = debug_rx.next().await {
200                let run_ix = current_run_ix.load(SeqCst);
201                let mut updated_example = updated_example.lock().unwrap();
202
203                let run_dir = if repetition_count > 1 {
204                    run_dir.join(format!("{:03}", run_ix))
205                } else {
206                    run_dir.clone()
207                };
208
209                match event {
210                    DebugEvent::EditPredictionStarted(request) => {
211                        assert_eq!(updated_example.predictions.len(), run_ix + 1);
212
213                        if let Some(prompt) = request.prompt {
214                            fs::write(run_dir.join("prediction_prompt.md"), &prompt)?;
215                            if matches!(provider, PredictionProvider::Zeta2(_)) {
216                                updated_example.prompt.get_or_insert(ExamplePrompt {
217                                    input: prompt,
218                                    expected_output: None,
219                                    rejected_output: None,
220                                    provider,
221                                    prefill: None,
222                                });
223                            }
224                        }
225                    }
226                    DebugEvent::EditPredictionFinished(request) => {
227                        assert_eq!(updated_example.predictions.len(), run_ix + 1);
228
229                        if let Some(output) = request.model_output {
230                            fs::write(run_dir.join("prediction_response.md"), &output)?;
231                            updated_example
232                                .predictions
233                                .last_mut()
234                                .unwrap()
235                                .actual_output = output;
236                        }
237                        if run_ix >= repetition_count {
238                            break;
239                        }
240                    }
241                    _ => {}
242                }
243            }
244            anyhow::Ok(())
245        }
246    });
247
248    for ix in 0..repetition_count {
249        current_run_ix.store(ix, SeqCst);
250        let run_dir = if repetition_count > 1 {
251            run_dir.join(format!("{:03}", ix))
252        } else {
253            run_dir.clone()
254        };
255
256        if repetition_count > 1 {
257            step_progress.set_substatus(format!(
258                "running prediction {}/{}",
259                ix + 1,
260                repetition_count
261            ));
262        } else {
263            step_progress.set_substatus("running prediction");
264        }
265
266        fs::create_dir_all(&run_dir)?;
267        if LATEST_EXAMPLE_RUN_DIR.is_symlink() {
268            fs::remove_file(&*LATEST_EXAMPLE_RUN_DIR)?;
269        }
270        #[cfg(unix)]
271        std::os::unix::fs::symlink(&run_dir, &*LATEST_EXAMPLE_RUN_DIR)?;
272        #[cfg(windows)]
273        std::os::windows::fs::symlink_dir(&run_dir, &*LATEST_EXAMPLE_RUN_DIR)?;
274
275        updated_example
276            .lock()
277            .unwrap()
278            .predictions
279            .push(ExamplePrediction {
280                actual_patch: None,
281                actual_output: String::new(),
282                actual_cursor: None,
283                error: None,
284                provider,
285                cumulative_logprob: None,
286                avg_logprob: None,
287            });
288
289        step_progress.set_substatus("requesting prediction");
290        let prediction = ep_store
291            .update(&mut cx, |store, cx| {
292                store.request_prediction(
293                    &state.project,
294                    &state.buffer,
295                    state.cursor_position,
296                    cloud_llm_client::PredictEditsRequestTrigger::Cli,
297                    cx,
298                )
299            })
300            .await?;
301
302        let actual_patch = prediction.and_then(|result| {
303            let prediction = result.prediction;
304            prediction
305                .edit_preview
306                .as_unified_diff(prediction.snapshot.file(), &prediction.edits)
307        });
308
309        let has_prediction = actual_patch.as_ref().is_some_and(|p| !p.is_empty());
310
311        updated_example
312            .lock()
313            .unwrap()
314            .predictions
315            .last_mut()
316            .unwrap()
317            .actual_patch = actual_patch;
318
319        if ix == repetition_count - 1 {
320            let (info, style) = if has_prediction {
321                ("predicted", InfoStyle::Normal)
322            } else {
323                ("no prediction", InfoStyle::Warning)
324            };
325            step_progress.set_info(info, style);
326        }
327    }
328
329    ep_store.update(&mut cx, |store, _| {
330        store.remove_project(&state.project);
331    });
332    debug_task.await?;
333
334    *example = Arc::into_inner(updated_example)
335        .ok_or_else(|| anyhow::anyhow!("Failed to unwrap Arc"))?
336        .into_inner()
337        .map_err(|_| anyhow::anyhow!("Failed to unwrap Mutex"))?;
338    Ok(())
339}
340
341async fn predict_teacher(
342    example: &mut Example,
343    backend: TeacherBackend,
344    batched: bool,
345    repetition_count: usize,
346    cache_only: bool,
347    step_progress: &crate::progress::StepProgress,
348) -> anyhow::Result<()> {
349    match backend {
350        TeacherBackend::Sonnet45 | TeacherBackend::Sonnet46 => {
351            predict_anthropic(
352                example,
353                backend,
354                batched,
355                repetition_count,
356                cache_only,
357                step_progress,
358            )
359            .await
360        }
361        TeacherBackend::Gpt52 | TeacherBackend::Gpt54 | TeacherBackend::Gpt55 => {
362            predict_openai(
363                example,
364                backend,
365                batched,
366                repetition_count,
367                cache_only,
368                step_progress,
369            )
370            .await
371        }
372    }
373}
374
375async fn predict_anthropic(
376    example: &mut Example,
377    backend: TeacherBackend,
378    batched: bool,
379    repetition_count: usize,
380    cache_only: bool,
381    step_progress: &crate::progress::StepProgress,
382) -> anyhow::Result<()> {
383    let llm_model_name = backend.model_name();
384    let max_tokens = 16384;
385    let llm_client = ANTHROPIC_CLIENT.get_or_init(|| {
386        let client = if batched {
387            AnthropicClient::batch(&crate::paths::LLM_CACHE_DB)
388        } else {
389            AnthropicClient::plain()
390        };
391        client.expect("Failed to create Anthropic client")
392    });
393
394    let prompt = example.prompt.as_ref().context("Prompt is required")?;
395
396    for ix in 0..repetition_count {
397        if repetition_count > 1 {
398            step_progress.set_substatus(format!(
399                "running prediction {}/{}",
400                ix + 1,
401                repetition_count
402            ));
403        } else {
404            step_progress.set_substatus("running prediction");
405        }
406
407        let messages = vec![anthropic::Message {
408            role: anthropic::Role::User,
409            content: vec![anthropic::RequestContent::Text {
410                text: prompt.input.clone(),
411                cache_control: None,
412            }],
413        }];
414
415        let seed = if repetition_count > 1 { Some(ix) } else { None };
416        let Some(response) = llm_client
417            .generate(llm_model_name, max_tokens, messages, seed, cache_only)
418            .await?
419        else {
420            // Request stashed for batched processing
421            continue;
422        };
423
424        let actual_output = response
425            .content
426            .into_iter()
427            .filter_map(|content| match content {
428                anthropic::ResponseContent::Text { text } => Some(text),
429                _ => None,
430            })
431            .collect::<Vec<String>>()
432            .join("\n");
433
434        let parser_provider = if batched {
435            example
436                .prompt
437                .as_ref()
438                .map(|prompt| prompt.provider)
439                .unwrap_or(PredictionProvider::Teacher(backend, ZetaFormat::default()))
440        } else {
441            match example.prompt.as_ref().map(|prompt| prompt.provider) {
442                Some(PredictionProvider::TeacherJumps(_))
443                | Some(PredictionProvider::TeacherJumpsNonBatching(_)) => {
444                    PredictionProvider::TeacherJumpsNonBatching(backend)
445                }
446                _ => PredictionProvider::TeacherNonBatching(backend, ZetaFormat::default()),
447            }
448        };
449
450        let parse_result = match parser_provider {
451            PredictionProvider::TeacherJumps(_)
452            | PredictionProvider::TeacherJumpsNonBatching(_) => {
453                TeacherJumpsPrompt::parse(example, &actual_output)
454            }
455            _ => TeacherPrompt::parse(example, &actual_output),
456        };
457        // A teacher response can parse as text yet describe an invalid edit
458        // (e.g. an edit span crossing non-contiguous snippets, or a truncated
459        // span). Record the rejection on the prediction instead of propagating
460        // it: the raw output is preserved for `parse-output`/inspection, and a
461        // single bad example no longer aborts an entire (already paid for) batch.
462        let (actual_patch, actual_cursor, error) = match parse_result {
463            Ok((patch, cursor)) => (Some(patch), cursor, None),
464            Err(err) => (None, None, Some(format!("{err:#}"))),
465        };
466
467        let prediction = ExamplePrediction {
468            actual_patch,
469            actual_output,
470            actual_cursor,
471            error,
472            provider: if batched {
473                match example.prompt.as_ref().map(|prompt| prompt.provider) {
474                    Some(PredictionProvider::TeacherJumps(_)) => {
475                        PredictionProvider::TeacherJumps(backend)
476                    }
477                    _ => PredictionProvider::Teacher(backend, ZetaFormat::default()),
478                }
479            } else {
480                match example.prompt.as_ref().map(|prompt| prompt.provider) {
481                    Some(PredictionProvider::TeacherJumps(_))
482                    | Some(PredictionProvider::TeacherJumpsNonBatching(_)) => {
483                        PredictionProvider::TeacherJumpsNonBatching(backend)
484                    }
485                    _ => PredictionProvider::TeacherNonBatching(backend, ZetaFormat::default()),
486                }
487            },
488            cumulative_logprob: None,
489            avg_logprob: None,
490        };
491
492        example.predictions.push(prediction);
493    }
494    Ok(())
495}
496
497async fn predict_openai(
498    example: &mut Example,
499    backend: TeacherBackend,
500    batched: bool,
501    repetition_count: usize,
502    cache_only: bool,
503    step_progress: &crate::progress::StepProgress,
504) -> anyhow::Result<()> {
505    let llm_model_name = backend.model_name();
506    let max_tokens = 16384;
507    let llm_client = OPENAI_CLIENT.get_or_init(|| {
508        let client = if batched {
509            OpenAiClient::batch(&crate::paths::LLM_CACHE_DB)
510        } else {
511            OpenAiClient::plain()
512        };
513        client.expect("Failed to create OpenAI client")
514    });
515
516    let prompt = example.prompt.as_ref().context("Prompt is required")?;
517
518    for ix in 0..repetition_count {
519        if repetition_count > 1 {
520            step_progress.set_substatus(format!(
521                "running prediction {}/{}",
522                ix + 1,
523                repetition_count
524            ));
525        } else {
526            step_progress.set_substatus("running prediction");
527        }
528
529        let messages = vec![open_ai::RequestMessage::User {
530            content: open_ai::MessageContent::Plain(prompt.input.clone()),
531        }];
532
533        let seed = if repetition_count > 1 { Some(ix) } else { None };
534        let Some(response) = llm_client
535            .generate(llm_model_name, max_tokens, messages, seed, cache_only)
536            .await?
537        else {
538            // Request stashed for batched processing
539            continue;
540        };
541
542        let actual_output = response
543            .choices
544            .into_iter()
545            .filter_map(|choice| match choice.message {
546                open_ai::RequestMessage::Assistant { content, .. } => content.map(|c| match c {
547                    open_ai::MessageContent::Plain(text) => text,
548                    open_ai::MessageContent::Multipart(parts) => parts
549                        .into_iter()
550                        .filter_map(|p| match p {
551                            open_ai::MessagePart::Text { text } => Some(text),
552                            _ => None,
553                        })
554                        .collect::<Vec<_>>()
555                        .concat(),
556                }),
557                _ => None,
558            })
559            .collect::<Vec<String>>()
560            .join("\n");
561
562        let parser_provider = if batched {
563            example
564                .prompt
565                .as_ref()
566                .map(|prompt| prompt.provider)
567                .unwrap_or(PredictionProvider::Teacher(backend, ZetaFormat::default()))
568        } else {
569            match example.prompt.as_ref().map(|prompt| prompt.provider) {
570                Some(PredictionProvider::TeacherJumps(_))
571                | Some(PredictionProvider::TeacherJumpsNonBatching(_)) => {
572                    PredictionProvider::TeacherJumpsNonBatching(backend)
573                }
574                _ => PredictionProvider::TeacherNonBatching(backend, ZetaFormat::default()),
575            }
576        };
577
578        let parse_result = match parser_provider {
579            PredictionProvider::TeacherJumps(_)
580            | PredictionProvider::TeacherJumpsNonBatching(_) => {
581                TeacherJumpsPrompt::parse(example, &actual_output)
582            }
583            _ => TeacherPrompt::parse(example, &actual_output),
584        };
585        // See `predict_anthropic`: an unparseable/invalid teacher edit is
586        // recorded as a per-prediction error rather than aborting the batch.
587        let (actual_patch, actual_cursor, error) = match parse_result {
588            Ok((patch, cursor)) => (Some(patch), cursor, None),
589            Err(err) => (None, None, Some(format!("{err:#}"))),
590        };
591
592        let prediction = ExamplePrediction {
593            actual_patch,
594            actual_output,
595            actual_cursor,
596            error,
597            provider: if batched {
598                match example.prompt.as_ref().map(|prompt| prompt.provider) {
599                    Some(PredictionProvider::TeacherJumps(_)) => {
600                        PredictionProvider::TeacherJumps(backend)
601                    }
602                    _ => PredictionProvider::Teacher(backend, ZetaFormat::default()),
603                }
604            } else {
605                match example.prompt.as_ref().map(|prompt| prompt.provider) {
606                    Some(PredictionProvider::TeacherJumps(_))
607                    | Some(PredictionProvider::TeacherJumpsNonBatching(_)) => {
608                        PredictionProvider::TeacherJumpsNonBatching(backend)
609                    }
610                    _ => PredictionProvider::TeacherNonBatching(backend, ZetaFormat::default()),
611                }
612            },
613            cumulative_logprob: None,
614            avg_logprob: None,
615        };
616
617        example.predictions.push(prediction);
618    }
619    Ok(())
620}
621
622pub async fn predict_baseten(
623    example: &mut Example,
624    format: ZetaFormat,
625    step_progress: &StepProgress,
626) -> anyhow::Result<()> {
627    let model_id =
628        std::env::var("ZED_ZETA_MODEL").context("ZED_ZETA_MODEL environment variable required")?;
629
630    let api_key =
631        std::env::var("BASETEN_API_KEY").context("BASETEN_API_KEY environment variable not set")?;
632
633    let prompt = example.prompt.as_ref().context("Prompt is required")?;
634    let prompt_text = prompt.input.clone();
635    let prefill = prompt.prefill.clone().unwrap_or_default();
636
637    step_progress.set_substatus("running prediction via baseten");
638
639    let environment: String = <&'static str>::from(&format).to_lowercase();
640    let url = format!(
641        "https://model-{model_id}.api.baseten.co/environments/{environment}/sync/v1/completions"
642    );
643
644    let request_body = RawCompletionRequest {
645        model: model_id,
646        prompt: prompt_text.clone(),
647        max_tokens: Some(2048),
648        temperature: Some(0.),
649        stop: vec![],
650        environment: None,
651    };
652
653    let body_bytes =
654        serde_json::to_vec(&request_body).context("Failed to serialize request body")?;
655
656    let http_client: Arc<dyn HttpClient> = Arc::new(ReqwestClient::new());
657    let request = http_client::Request::builder()
658        .method(Method::POST)
659        .uri(&url)
660        .header("Content-Type", "application/json")
661        .header("Authorization", format!("Api-Key {api_key}"))
662        .body(AsyncBody::from(body_bytes))?;
663
664    let mut response = http_client.send(request).await?;
665    let status = response.status();
666
667    let mut body = String::new();
668    response
669        .body_mut()
670        .read_to_string(&mut body)
671        .await
672        .context("Failed to read Baseten response body")?;
673
674    if !status.is_success() {
675        anyhow::bail!("Baseten API returned {status}: {body}");
676    }
677
678    let completion: RawCompletionResponse =
679        serde_json::from_str(&body).context("Failed to parse Baseten response")?;
680
681    let actual_output = completion
682        .choices
683        .into_iter()
684        .next()
685        .map(|choice| choice.text)
686        .unwrap_or_default();
687
688    let actual_output = format!("{prefill}{actual_output}");
689
690    let (actual_patch, actual_cursor) =
691        parse_prediction_output(example, &actual_output, PredictionProvider::Zeta2(format))?;
692
693    let prediction = ExamplePrediction {
694        actual_patch: Some(actual_patch),
695        actual_output,
696        actual_cursor,
697        error: None,
698        provider: PredictionProvider::Baseten(format),
699        cumulative_logprob: None,
700        avg_logprob: None,
701    };
702
703    example.predictions.push(prediction);
704    Ok(())
705}
706
707pub async fn sync_batches(provider: Option<&PredictionProvider>) -> anyhow::Result<()> {
708    match provider {
709        Some(PredictionProvider::Teacher(backend, _))
710        | Some(PredictionProvider::TeacherJumps(backend)) => match backend {
711            TeacherBackend::Sonnet45 | TeacherBackend::Sonnet46 => {
712                let llm_client = ANTHROPIC_CLIENT.get_or_init(|| {
713                    AnthropicClient::batch(&crate::paths::LLM_CACHE_DB)
714                        .expect("Failed to create Anthropic client")
715                });
716                llm_client
717                    .sync_batches()
718                    .await
719                    .context("Failed to sync Anthropic batches")?;
720            }
721            TeacherBackend::Gpt52 | TeacherBackend::Gpt54 | TeacherBackend::Gpt55 => {
722                let llm_client = OPENAI_CLIENT.get_or_init(|| {
723                    OpenAiClient::batch(&crate::paths::LLM_CACHE_DB)
724                        .expect("Failed to create OpenAI client")
725                });
726                llm_client
727                    .sync_batches()
728                    .await
729                    .context("Failed to sync OpenAI batches")?;
730            }
731        },
732        _ => (),
733    };
734    Ok(())
735}
736
737pub async fn reprocess_after_batch_wait(
738    examples: &mut [Example],
739    args: &PredictArgs,
740) -> anyhow::Result<()> {
741    let (Some(PredictionProvider::Teacher(backend, _))
742    | Some(PredictionProvider::TeacherJumps(backend))) = args.provider
743    else {
744        return Ok(());
745    };
746
747    let mut reprocessed = 0;
748    for example in examples.iter_mut() {
749        let has_prediction = example
750            .predictions
751            .iter()
752            .any(|p| p.actual_patch.is_some() || !p.actual_output.is_empty());
753        if has_prediction || example.prompt.is_none() {
754            continue;
755        }
756
757        let example_progress = Progress::global().start_group(&example.spec.name);
758        let step_progress = example_progress.start(Step::Predict);
759        predict_teacher(
760            example,
761            backend,
762            true,
763            args.repetitions,
764            false,
765            &step_progress,
766        )
767        .await?;
768        reprocessed += 1;
769    }
770
771    if reprocessed > 0 {
772        eprintln!("Reprocessed {} example(s) with batch results", reprocessed);
773    }
774
775    Ok(())
776}
777
778pub async fn wait_for_batches(provider: Option<&PredictionProvider>) -> anyhow::Result<()> {
779    let poll_interval = std::time::Duration::from_secs(30);
780
781    loop {
782        let pending = pending_batch_count(provider)?;
783        if pending == 0 {
784            break;
785        }
786
787        eprintln!(
788            "Waiting for {} pending batch request(s) to complete... (polling every {}s)",
789            pending,
790            poll_interval.as_secs()
791        );
792        std::thread::sleep(poll_interval);
793
794        sync_batches(provider).await?;
795    }
796
797    Ok(())
798}
799
800fn pending_batch_count(provider: Option<&PredictionProvider>) -> anyhow::Result<usize> {
801    match provider {
802        Some(PredictionProvider::Teacher(backend, _))
803        | Some(PredictionProvider::TeacherJumps(backend)) => match backend {
804            TeacherBackend::Sonnet45 | TeacherBackend::Sonnet46 => {
805                let llm_client = ANTHROPIC_CLIENT.get_or_init(|| {
806                    AnthropicClient::batch(&crate::paths::LLM_CACHE_DB)
807                        .expect("Failed to create Anthropic client")
808                });
809                llm_client.pending_batch_count()
810            }
811            TeacherBackend::Gpt52 | TeacherBackend::Gpt54 | TeacherBackend::Gpt55 => {
812                let llm_client = OPENAI_CLIENT.get_or_init(|| {
813                    OpenAiClient::batch(&crate::paths::LLM_CACHE_DB)
814                        .expect("Failed to create OpenAI client")
815                });
816                llm_client.pending_batch_count()
817            }
818        },
819        _ => Ok(0),
820    }
821}
822
Served at tenant.openagents/omega Member data and write actions are omitted.