Skip to repository content822 lines · 28.9 KB · rust
tenant.openagents/omega
No repository description is available.
OpenAgents Git authority 2026-07-28T03:54:16.929Z Public web read
NIP-34 coordinate
30617:7649603503856e5148d571eac2766b288a8ff1e9e35d380337a1d2b0015b4f92:omegaMaintainersHidden in public view
References2 branches · 1 tag
Read-only clone
git clone https://openagents.com/git/tenant.openagents/omega.gitBrowse files
predict.rs
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