Skip to repository content1707 lines · 65.7 KB · rust
tenant.openagents/omega
No repository description is available.
OpenAgents Git authority 2026-07-28T02:53:44.123Z 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
main.rs
1mod anthropic_client;
2mod distill;
3mod example;
4mod filter_languages;
5mod format_prompt;
6mod git;
7mod headless;
8
9mod load_project;
10mod metrics;
11mod openai_client;
12mod parse_output;
13mod paths;
14mod predict;
15mod progress;
16mod prompt_assets;
17mod pull_examples;
18mod qa;
19mod reorder_patch;
20mod repair;
21mod retrieve_context;
22mod score;
23mod split_commit;
24mod split_dataset;
25
26mod synthesize;
27mod truncate_expected_patch;
28mod word_diff;
29use anyhow::Context as _;
30use clap::{Args, CommandFactory, Parser, Subcommand, ValueEnum};
31use collections::{HashMap, HashSet};
32use edit_prediction::EditPredictionStore;
33use futures::channel::mpsc;
34use futures::{SinkExt as _, StreamExt as _};
35use gaoya::minhash::{
36 MinHashIndex, MinHasher, MinHasher32, calculate_minhash_params, compute_minhash_similarity,
37};
38use gpui::{AppContext as _, BackgroundExecutor, Task};
39use zeta_prompt::{ContextSource, ZetaFormat};
40
41use reqwest_client::ReqwestClient;
42use serde::{Deserialize, Deserializer, Serialize, Serializer};
43use std::collections::VecDeque;
44use std::env;
45use std::fmt::Display;
46use std::fs::{File, OpenOptions};
47use std::hash::{Hash, Hasher};
48use std::io::{BufRead, BufReader, BufWriter, Write};
49use std::str::FromStr;
50use std::sync::Mutex;
51use std::{path::PathBuf, sync::Arc};
52
53use crate::distill::run_distill;
54use crate::example::{Example, group_examples_by_repo, read_example_files};
55use crate::filter_languages::{FilterLanguagesArgs, run_filter_languages};
56use crate::format_prompt::run_format_prompt;
57use crate::load_project::run_load_project;
58use crate::paths::{FAILED_EXAMPLES_DIR, RUN_DIR};
59use crate::predict::run_prediction;
60use crate::progress::Progress;
61use crate::pull_examples::{fetch_settled_examples_after, parse_settled_after_input};
62use crate::retrieve_context::{
63 ContextRetrievalType, context_sources_for_types, run_context_retrieval,
64};
65use crate::score::run_scoring;
66use crate::split_commit::SplitCommitArgs;
67use crate::split_dataset::SplitArgs;
68use crate::synthesize::{SynthesizeConfig, run_synthesize};
69use crate::truncate_expected_patch::TruncatePatchArgs;
70
71#[derive(Parser, Debug)]
72#[command(name = "ep")]
73struct EpArgs {
74 #[arg(long, default_value_t = false)]
75 printenv: bool,
76 #[clap(long, default_value_t = 10, global = true)]
77 max_parallelism: usize,
78 /// Process all examples from a repository together instead of distributing examples across workers.
79 #[clap(long, default_value_t = false, global = true)]
80 group_by_repo: bool,
81 /// The limit for the number of examples to process
82 /// Default is unlimited for processing local datasets, 5000 when pulling from snowflake
83 #[clap(long, global = true)]
84 limit: Option<usize>,
85 #[clap(long, global = true)]
86 offset: Option<usize>,
87 /// Filter examples by name
88 #[clap(long, global = true)]
89 name: Option<String>,
90 /// Filter examples by repository
91 #[clap(long, global = true)]
92 repo: Option<String>,
93 /// Deduplicate by cursor position and keep at most this many examples per cluster
94 #[clap(long, global = true)]
95 max_duplicates: Option<usize>,
96 #[command(subcommand)]
97 command: Option<Command>,
98 /// Input file paths
99 #[clap(global = true)]
100 inputs: Vec<PathBuf>,
101 #[arg(long, short, global = true)]
102 output: Option<PathBuf>,
103 #[arg(long, short, global = true)]
104 in_place: bool,
105 #[arg(long, global = true)]
106 failfast: bool,
107 /// How to handle failed examples in output: keep them or skip them.
108 /// Failed examples are always logged to the run's failed directory.
109 #[arg(long, global = true, default_value = "keep")]
110 failed: FailedHandling,
111 /// Output as markdown files instead of JSONL. When set, -o specifies a directory
112 /// where one .md file per example will be written (named after each example).
113 #[arg(long, short, global = true)]
114 markdown: bool,
115}
116
117/// Controls whether failed examples are included in the main output.
118/// Failed examples are always logged to the run's failed/ directory regardless of this setting.
119#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, ValueEnum)]
120pub enum FailedHandling {
121 /// Include failed examples in the main output (default)
122 #[default]
123 Keep,
124 /// Exclude failed examples from the main output
125 Skip,
126 /// Skip writing files
127 SkipNoFiles,
128}
129
130#[derive(Args, Debug, Clone)]
131struct ContextArgs {
132 /// Which context collectors to run.
133 /// May be repeated or comma-delimited, e.g. `--type=all,oracle-file`.
134 #[arg(long = "type", value_enum, value_delimiter = ',')]
135 context_types: Vec<ContextRetrievalType>,
136 /// Recompute context even if the example already has related files.
137 #[arg(long, short = 'f', default_value_t = false)]
138 force: bool,
139}
140
141impl ContextArgs {
142 fn context_types(&self) -> Vec<ContextRetrievalType> {
143 if self.context_types.is_empty() {
144 vec![ContextRetrievalType::Lsp]
145 } else {
146 self.context_types.clone()
147 }
148 }
149}
150
151const INPUTS_HELP: &str = r#"
152Inputs can be file paths or special specifiers:
153
154 path
155 Path to an example(s) file (.md, .json, or .jsonl)
156
157 captured-after:{timestamp}
158 Fetch captured examples from Snowflake after the given RFC3339 timestamp.
159 These are examples captured via the "Capture Edit Prediction Example" action.
160
161 rejected-after:{timestamp}
162 Fetch rejected edit predictions from Snowflake after the given RFC3339 timestamp.
163 These are predictions that were shown to users but rejected (useful for DPO training).
164
165 settled-after:{timestamp}
166 Fetch settled stream examples from Snowflake after the given RFC3339 timestamp.
167 These are examples from the edit prediction settled stream.
168
169 rated-after:{timestamp}
170 Fetch user-rated edit predictions from Snowflake after the given RFC3339 timestamp.
171 These are predictions that users explicitly rated as positive or negative via the
172 rate completions modal. Only zeta2 predictions are included.
173 - Positive ratings: output becomes expected_patches
174 - Negative ratings: output becomes rejected_patch
175
176 rated-positive-after:{timestamp}
177 Same as rated-after, but only fetches positively rated predictions.
178
179 rated-negative-after:{timestamp}
180 Same as rated-after, but only fetches negatively rated predictions.
181
182 Required environment variables to connect to Snowflake:
183 EP_SNOWFLAKE_API_KEY
184 EP_SNOWFLAKE_BASE_URL
185
186 Optional:
187 EP_SNOWFLAKE_ROLE
188
189Examples:
190
191 # Read examples from a file
192 ep read examples.jsonl -o output.jsonl
193
194 # Read captured examples after a timestamp
195 ep read captured-after:2025-01-01T00:00:00Z -o captured.jsonl
196
197 # Read rejected predictions for DPO training
198 ep read rejected-after:2025-01-01T00:00:00Z -o rejected.jsonl
199
200 # Read user-rated predictions
201 ep read rated-after:2025-01-01T00:00:00Z -o rated.jsonl
202
203 # Read settled stream examples
204 ep read settled-after:2025-01-01T00:00:00Z -o settled.jsonl
205
206 # Read only positively rated predictions
207 ep read rated-positive-after:2025-01-01T00:00:00Z -o positive.jsonl
208
209 # Read only negatively rated predictions
210 ep read rated-negative-after:2025-01-01T00:00:00Z -o negative.jsonl
211
212 # Mix multiple input sources
213 ep predict examples.jsonl captured-after:2025-01-01T00:00:00Z
214"#;
215
216#[derive(Subcommand, Debug, Clone)]
217enum Command {
218 /// Read examples from files or fetch from Snowflake, output as .jsonl
219 Read(ReadArgs),
220 /// Create git worktrees for each example and load file contents
221 LoadProject,
222 /// Retrieve context for input examples.
223 Context(ContextArgs),
224 /// Generate a prompt string for a specific model
225 FormatPrompt(FormatPromptArgs),
226 /// Runs edit prediction
227 Predict(PredictArgs),
228 /// Parse model outputs (actual_output) into unified diffs (actual_patch).
229 /// Requires format-prompt to have been run first. Uses provider from prompt.
230 ParseOutput,
231 /// Computes a score based on actual and expected patches
232 Score(PredictArgs),
233 /// Prepares a distillation dataset by copying expected outputs to
234 /// predicted outputs and removing actual outputs and prompts.
235 Distill,
236 /// Print aggregated scores
237 Eval(EvalArgs),
238 /// Generate eval examples by analyzing git commits from a repository
239 Synthesize(SynthesizeArgs),
240 /// Remove git repositories and worktrees
241 Clean,
242 /// Generate an evaluation example by splitting a chronologically-ordered commit
243 SplitCommit(SplitCommitArgs),
244 /// Truncate expected patch by the given criteria
245 TruncatePatch(TruncatePatchArgs),
246 /// Split a JSONL dataset into multiple files (stratified by repository_url if present)
247 Split(SplitArgs),
248 /// Filter a JSONL dataset by programming language (based on cursor_path extension)
249 FilterLanguages(FilterLanguagesArgs),
250 /// Import Anthropic batch results by batch IDs (useful for recovering after database loss)
251 ImportBatch(ImportBatchArgs),
252 /// Assess the quality of predictions using LLM-as-a-judge
253 Qa(qa::QaArgs),
254 /// Repair predictions that received poor QA scores by generating improved predictions
255 Repair(repair::RepairArgs),
256 /// Print all valid zeta formats (lowercase, one per line)
257 PrintZetaFormats,
258}
259
260impl Display for Command {
261 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
262 match self {
263 Command::Read(_) => write!(f, "read"),
264 Command::LoadProject => write!(f, "load-project"),
265 Command::Context(args) => {
266 write!(f, "context --type=")?;
267 for (index, context_type) in args.context_types().iter().enumerate() {
268 if index > 0 {
269 write!(f, ",")?;
270 }
271 write!(f, "{}", context_type)?;
272 }
273 if args.force {
274 write!(f, " --force")?;
275 }
276 Ok(())
277 }
278 Command::FormatPrompt(args) => {
279 write!(f, "format-prompt --provider={}", args.provider)
280 }
281 Command::Predict(args) => match &args.provider {
282 Some(provider) => write!(f, "predict --provider={}", provider),
283 None => write!(f, "predict"),
284 },
285 Command::ParseOutput => write!(f, "parse-output"),
286 Command::Score(args) => match &args.provider {
287 Some(provider) => write!(f, "score --provider={}", provider),
288 None => write!(f, "score"),
289 },
290 Command::Distill => write!(f, "distill"),
291 Command::Eval(args) => {
292 write!(f, "eval")?;
293 if args.context_only {
294 write!(f, " --context-only")?;
295 }
296 if !args.context_types.is_empty() {
297 write!(f, " --type=")?;
298 for (index, context_type) in args.context_types.iter().enumerate() {
299 if index > 0 {
300 write!(f, ",")?;
301 }
302 write!(f, "{}", context_type)?;
303 }
304 }
305 if args.related_context_limit != score::EVAL_RELATED_CONTEXT_TOKENS_LIMIT {
306 write!(f, " --related-context-limit={}", args.related_context_limit)?;
307 }
308 if let Some(provider) = &args.predict.provider {
309 write!(f, " --provider={}", provider)?;
310 }
311 Ok(())
312 }
313 Command::Synthesize(args) => {
314 write!(f, "synthesize --repos {}", args.repos.join(" "))
315 }
316 Command::Clean => write!(f, "clean"),
317 Command::SplitCommit(_) => write!(f, "split-commit"),
318 Command::TruncatePatch(_) => write!(f, "truncate-patch"),
319 Command::Split(_) => write!(f, "split"),
320 Command::FilterLanguages(_) => write!(f, "filter-languages"),
321 Command::ImportBatch(args) => {
322 write!(f, "import-batch --batch-ids {}", args.batch_ids.join(" "))
323 }
324 Command::Qa(_) => {
325 write!(f, "qa")
326 }
327 Command::Repair(_) => {
328 write!(f, "repair")
329 }
330 Command::PrintZetaFormats => {
331 write!(f, "print-zeta-formats")
332 }
333 }
334 }
335}
336
337#[derive(Debug, Args, Clone)]
338#[command(after_help = INPUTS_HELP)]
339struct ReadArgs {}
340
341#[derive(Debug, Args, Clone)]
342struct FormatPromptArgs {
343 #[clap(long, short('p'), default_value_t = PredictionProvider::default())]
344 provider: PredictionProvider,
345 /// Token budget for related-file context in teacher-jumps prompts.
346 #[clap(long, default_value_t = format_prompt::TeacherJumpsPrompt::DEFAULT_RELATED_FILES_BUDGET)]
347 related_files_budget: usize,
348}
349
350#[derive(Debug, Args, Clone)]
351struct PredictArgs {
352 #[clap(long, short('p'))]
353 provider: Option<PredictionProvider>,
354 #[clap(long, default_value_t = 1)]
355 repetitions: usize,
356 /// Only use cached responses, don't queue new requests for batching
357 #[clap(long)]
358 cache_only: bool,
359 /// Wait for all batches to complete before exiting (only applies to batched providers like teacher)
360 #[clap(long)]
361 wait: bool,
362}
363
364#[derive(Debug, Args, Clone)]
365struct EvalArgs {
366 #[clap(flatten)]
367 predict: PredictArgs,
368 /// Only compute editable context coverage from expected patches and retrieved context.
369 #[clap(long)]
370 context_only: bool,
371 /// Only score persisted related context excerpts from these context types.
372 /// May be repeated or comma-delimited, e.g. `--type=current-file,edit-history`.
373 #[arg(long = "type", value_enum, value_delimiter = ',')]
374 context_types: Vec<ContextRetrievalType>,
375 /// Maximum number of retrieved context tokens to include when scoring.
376 #[clap(long, default_value_t = score::EVAL_RELATED_CONTEXT_TOKENS_LIMIT)]
377 related_context_limit: usize,
378 /// Path to write summary scores as JSON
379 #[clap(long)]
380 summary_json: Option<PathBuf>,
381 /// Print all individual example lines (default: up to 20)
382 #[clap(long)]
383 verbose: bool,
384}
385
386impl EvalArgs {
387 fn context_source_filter(&self) -> Option<Vec<ContextSource>> {
388 if self.context_types.is_empty() {
389 None
390 } else {
391 Some(context_sources_for_types(&self.context_types))
392 }
393 }
394}
395
396#[derive(Clone, Copy, Default, Debug, PartialEq, Eq, Hash)]
397pub enum TeacherBackend {
398 Sonnet46,
399 #[default]
400 Sonnet45,
401 Gpt52,
402 Gpt54,
403 Gpt55,
404}
405
406impl std::fmt::Display for TeacherBackend {
407 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
408 match self {
409 TeacherBackend::Sonnet46 => write!(f, "sonnet46"),
410 TeacherBackend::Sonnet45 => write!(f, "sonnet45"),
411 TeacherBackend::Gpt52 => write!(f, "gpt52"),
412 TeacherBackend::Gpt54 => write!(f, "gpt54"),
413 TeacherBackend::Gpt55 => write!(f, "gpt55"),
414 }
415 }
416}
417
418impl std::str::FromStr for TeacherBackend {
419 type Err = anyhow::Error;
420
421 fn from_str(s: &str) -> Result<Self, Self::Err> {
422 match s.to_lowercase().as_str() {
423 "sonnet45" | "sonnet" | "claude" => Ok(TeacherBackend::Sonnet45),
424 "sonnet46" => Ok(TeacherBackend::Sonnet46),
425 "gpt52" => Ok(TeacherBackend::Gpt52),
426 "gpt54" | "gpt" | "openai" => Ok(TeacherBackend::Gpt54),
427 "gpt55" => Ok(TeacherBackend::Gpt55),
428 "v0114180editableregion" => Ok(TeacherBackend::Sonnet45),
429 _ => anyhow::bail!(
430 "unknown teacher backend `{s}`. Valid options: sonnet45, sonnet46, gpt52, gpt54, gpt55"
431 ),
432 }
433 }
434}
435
436impl TeacherBackend {
437 pub fn model_name(&self) -> &'static str {
438 match self {
439 TeacherBackend::Sonnet45 => "claude-sonnet-4-5",
440 TeacherBackend::Sonnet46 => "claude-sonnet-4-6",
441 TeacherBackend::Gpt52 => "gpt-5.2",
442 TeacherBackend::Gpt54 => "gpt-5.4",
443 TeacherBackend::Gpt55 => "gpt-5.5",
444 }
445 }
446}
447
448#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
449enum PredictionProvider {
450 Mercury,
451 Zeta1,
452 Zeta2(ZetaFormat),
453 Baseten(ZetaFormat),
454 Teacher(TeacherBackend, ZetaFormat),
455 TeacherJumps(TeacherBackend),
456 TeacherNonBatching(TeacherBackend, ZetaFormat),
457 TeacherJumpsNonBatching(TeacherBackend),
458 Repair,
459}
460
461impl Default for PredictionProvider {
462 fn default() -> Self {
463 PredictionProvider::Zeta2(ZetaFormat::default())
464 }
465}
466
467impl std::fmt::Display for PredictionProvider {
468 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
469 match self {
470 PredictionProvider::Mercury => write!(f, "mercury"),
471 PredictionProvider::Zeta1 => write!(f, "zeta1"),
472 PredictionProvider::Zeta2(format) => write!(f, "zeta2:{format}"),
473 PredictionProvider::Baseten(format) => write!(f, "baseten:{format}"),
474 PredictionProvider::Teacher(backend, format) => {
475 write!(f, "teacher:{backend}:{format:?}")
476 }
477 PredictionProvider::TeacherJumps(backend) => {
478 write!(f, "teacher-jumps:{backend}")
479 }
480 PredictionProvider::TeacherNonBatching(backend, format) => {
481 write!(f, "teacher-non-batching:{backend}:{format:?}")
482 }
483 PredictionProvider::TeacherJumpsNonBatching(backend) => {
484 write!(f, "teacher-jumps-non-batching:{backend}")
485 }
486 PredictionProvider::Repair => write!(f, "repair"),
487 }
488 }
489}
490
491impl std::str::FromStr for PredictionProvider {
492 type Err = anyhow::Error;
493
494 fn from_str(s: &str) -> Result<Self, Self::Err> {
495 let (provider, arg) = s.split_once(':').map_or((s, None), |(p, a)| (p, Some(a)));
496
497 let provider_lower = provider.to_lowercase();
498 match provider_lower.as_str() {
499 "mercury" => Ok(PredictionProvider::Mercury),
500 "zeta1" => Ok(PredictionProvider::Zeta1),
501 "zeta2" => {
502 let format = arg.map(ZetaFormat::parse).transpose()?.unwrap_or_default();
503 Ok(PredictionProvider::Zeta2(format))
504 }
505 "teacher" => {
506 let (backend, format) = parse_teacher_args(arg)?;
507 Ok(PredictionProvider::Teacher(backend, format))
508 }
509 "teacher-non-batching" | "teacher_non_batching" => {
510 let (backend, format) = parse_teacher_args(arg)?;
511 Ok(PredictionProvider::TeacherNonBatching(backend, format))
512 }
513 "teacher-jumps" | "teacher_jumps" => {
514 let backend = arg
515 .map(|a| a.parse())
516 .transpose()?
517 .unwrap_or(TeacherBackend::default());
518 Ok(PredictionProvider::TeacherJumps(backend))
519 }
520 "teacher-jumps-non-batching" | "teacher_jumps_non_batching" => {
521 let backend = arg
522 .map(|a| a.parse())
523 .transpose()?
524 .unwrap_or(TeacherBackend::default());
525 Ok(PredictionProvider::TeacherJumpsNonBatching(backend))
526 }
527 "repair" => Ok(PredictionProvider::Repair),
528 "baseten" => {
529 let format = arg
530 .map(ZetaFormat::parse)
531 .transpose()?
532 .unwrap_or(ZetaFormat::default());
533 Ok(PredictionProvider::Baseten(format))
534 }
535 _ => {
536 anyhow::bail!(
537 "unknown provider `{provider}`. Valid options: mercury, zeta1, zeta2, zeta2:<version>, teacher, teacher:<backend>, teacher-jumps, teacher-jumps:<backend>, teacher-non-batching, teacher-jumps-non-batching, repair\n\
538 For zeta2, you can optionally specify a version like `zeta2:ordered` or `zeta2:V0113_Ordered`.\n\
539 For teacher providers, you can specify a backend like `teacher:sonnet46`, `teacher-jumps:sonnet46`, `teacher-jumps-non-batching:sonnet46`, or `teacher:gpt52`.\n\
540 Available zeta versions:\n{}",
541 ZetaFormat::options_as_string()
542 )
543 }
544 }
545 }
546}
547
548fn parse_teacher_args(arg: Option<&str>) -> Result<(TeacherBackend, ZetaFormat), anyhow::Error> {
549 let mut backend = TeacherBackend::default();
550 let mut format = ZetaFormat::default();
551
552 for arg in arg.unwrap_or_default().split(':') {
553 if arg.is_empty() {
554 continue;
555 }
556
557 if let Ok(parsed_backend) = TeacherBackend::from_str(arg) {
558 backend = parsed_backend;
559 } else if let Ok(parsed_format) = ZetaFormat::parse(arg) {
560 format = parsed_format;
561 } else {
562 anyhow::bail!("unknown teacher backend or zeta format `{arg}`");
563 }
564 }
565
566 Ok((backend, format))
567}
568
569impl Serialize for PredictionProvider {
570 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
571 where
572 S: Serializer,
573 {
574 serializer.serialize_str(&self.to_string())
575 }
576}
577
578impl<'de> Deserialize<'de> for PredictionProvider {
579 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
580 where
581 D: Deserializer<'de>,
582 {
583 let s = String::deserialize(deserializer)?;
584 s.parse().map_err(serde::de::Error::custom)
585 }
586}
587
588#[derive(Debug, Args, Clone)]
589struct SynthesizeArgs {
590 /// Repository URLs (git@github.com:owner/repo or https://...)
591 #[clap(long, required = true, num_args = 1..)]
592 repos: Vec<String>,
593
594 /// Number of examples to generate per repository
595 #[clap(long, default_value_t = 5)]
596 count: usize,
597
598 /// Maximum commits to scan per repository before giving up
599 #[clap(long, default_value_t = 100)]
600 max_commits: usize,
601
602 /// Ignore state file and reprocess all commits
603 #[clap(long)]
604 fresh: bool,
605}
606
607#[derive(Debug, Args, Clone)]
608struct ImportBatchArgs {
609 /// Batch IDs to import (e.g., msgbatch_xxx for Anthropic, batch_xxx for OpenAI)
610 #[clap(long, required = true, num_args = 1..)]
611 batch_ids: Vec<String>,
612 /// Which provider's batches to import (anthropic or openai)
613 #[clap(long, default_value = "anthropic")]
614 provider: BatchProvider,
615}
616
617#[derive(Debug, Clone, Copy, PartialEq, Eq, clap::ValueEnum)]
618enum BatchProvider {
619 Anthropic,
620 Openai,
621}
622
623#[cfg(test)]
624mod tests {
625 use super::*;
626
627 #[test]
628 fn prediction_provider_jumps_non_batched_round_trips_to_primary_spelling() {
629 let provider: PredictionProvider = "teacher-jumps-non-batching:sonnet46".parse().unwrap();
630 assert_eq!(
631 provider,
632 PredictionProvider::TeacherJumpsNonBatching(TeacherBackend::Sonnet46)
633 );
634 assert_eq!(provider.to_string(), "teacher-jumps-non-batching:sonnet46");
635 }
636
637 #[test]
638 fn prediction_provider_jumps_non_batched_alias_round_trips_to_primary_spelling() {
639 let provider: PredictionProvider = "teacher_jumps_non_batching:gpt52".parse().unwrap();
640 assert_eq!(
641 provider,
642 PredictionProvider::TeacherJumpsNonBatching(TeacherBackend::Gpt52)
643 );
644 assert_eq!(provider.to_string(), "teacher-jumps-non-batching:gpt52");
645 }
646}
647
648impl EpArgs {
649 fn output_path(&self) -> Option<PathBuf> {
650 if self.in_place {
651 if self.inputs.len() == 1 {
652 self.inputs.first().cloned()
653 } else {
654 panic!("--in-place requires exactly one input file")
655 }
656 } else {
657 self.output.clone()
658 }
659 }
660}
661
662/// Minimum Omega version required for Snowflake queries.
663/// This version introduced the current request schema with predicted edits in the edit
664/// history, and open source repos distinguished.
665const MIN_CAPTURE_VERSION: pull_examples::MinCaptureVersion = pull_examples::MinCaptureVersion {
666 major: 0,
667 minor: 224,
668 patch: 1,
669};
670
671fn deduplicate_examples(examples: &mut Vec<Example>, max_per_cluster: usize) {
672 let total_before_exact = examples.len();
673 let mut seen_positions = HashSet::default();
674 examples.retain(|example| seen_positions.insert(example.spec.cursor_position.clone()));
675 log::info!(
676 "exact duplicate filter: {total_before_exact} examples → {} examples ({} removed)",
677 examples.len(),
678 total_before_exact - examples.len(),
679 );
680
681 const JACCARD_THRESHOLD: f64 = 0.5;
682 const NUM_HASHES: usize = 128;
683 const TOKEN_NGRAM_SIZE: usize = 5;
684
685 let (num_bands, band_width) = calculate_minhash_params(JACCARD_THRESHOLD, NUM_HASHES);
686 let num_hashes = num_bands * band_width;
687 let minhasher = MinHasher32::new(num_hashes);
688 let mut index: MinHashIndex<u32, usize> =
689 MinHashIndex::new(num_bands, band_width, JACCARD_THRESHOLD);
690
691 let signatures: Vec<Vec<u32>> = examples
692 .iter()
693 .map(|example| {
694 let shingles = code_token_ngrams(&example.spec.cursor_position, TOKEN_NGRAM_SIZE);
695 minhasher.create_signature(shingles.iter())
696 })
697 .collect();
698
699 for (id, signature) in signatures.iter().enumerate() {
700 index.insert(id, signature.clone());
701 }
702
703 // Build clusters via union-find on LSH candidate pairs.
704 let mut parent: Vec<usize> = (0..examples.len()).collect();
705
706 fn find(parent: &mut Vec<usize>, mut x: usize) -> usize {
707 while parent[x] != x {
708 parent[x] = parent[parent[x]];
709 x = parent[x];
710 }
711 x
712 }
713
714 for (id, signature) in signatures.iter().enumerate() {
715 for candidate in index.query_owned(signature) {
716 let (a, b) = (find(&mut parent, id), find(&mut parent, candidate));
717 if a != b {
718 parent[a] = b;
719 }
720 }
721 }
722
723 let mut clusters: HashMap<usize, Vec<usize>> = HashMap::default();
724 for id in 0..examples.len() {
725 clusters.entry(find(&mut parent, id)).or_default().push(id);
726 }
727
728 let mut keep: HashSet<usize> = HashSet::default();
729 for members in clusters.values() {
730 let selected = greedy_max_min_diverse(members, &signatures, max_per_cluster);
731 keep.extend(selected);
732 }
733
734 let total = examples.len();
735 let mut kept_indices: Vec<usize> = keep.into_iter().collect();
736 kept_indices.sort();
737
738 let mut retained = Vec::with_capacity(kept_indices.len());
739 for index in kept_indices.into_iter().rev() {
740 retained.push(examples.swap_remove(index));
741 }
742 retained.reverse();
743
744 *examples = retained;
745 log::info!(
746 "near-duplicate filter: {total} examples → {} examples ({} removed)",
747 examples.len(),
748 total - examples.len(),
749 );
750}
751
752fn greedy_max_min_diverse(members: &[usize], signatures: &[Vec<u32>], k: usize) -> Vec<usize> {
753 if members.len() <= k {
754 return members.to_vec();
755 }
756
757 let mut selected = vec![members[0]];
758 let mut min_dist: HashMap<usize, f64> = HashMap::default();
759 for &member in &members[1..] {
760 let dist = 1.0 - compute_minhash_similarity(&signatures[selected[0]], &signatures[member]);
761 min_dist.insert(member, dist);
762 }
763
764 while selected.len() < k {
765 let &best = min_dist
766 .iter()
767 .max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))
768 .map(|(id, _)| id)
769 .expect("min_dist should not be empty when selected.len() < k");
770 selected.push(best);
771 min_dist.remove(&best);
772
773 let best_sig = &signatures[best];
774 for (member, current_min) in min_dist.iter_mut() {
775 let dist = 1.0 - compute_minhash_similarity(best_sig, &signatures[*member]);
776 if dist < *current_min {
777 *current_min = dist;
778 }
779 }
780 }
781
782 selected
783}
784
785fn code_token_ngrams(code: &str, ngram_size: usize) -> Vec<String> {
786 let tokens: Vec<&str> = word_diff::tokenize(code)
787 .into_iter()
788 .filter(|t| !t.trim().is_empty())
789 .collect();
790
791 if tokens.len() < ngram_size {
792 return vec![tokens.join("\0")];
793 }
794
795 tokens
796 .windows(ngram_size)
797 .map(|window| window.join("\0"))
798 .collect()
799}
800
801async fn load_examples(
802 http_client: Arc<dyn http_client::HttpClient>,
803 args: &EpArgs,
804 output_path: Option<&PathBuf>,
805 background_executor: BackgroundExecutor,
806) -> anyhow::Result<Vec<Example>> {
807 let mut captured_after_timestamps = Vec::new();
808 let mut rejected_after_timestamps = Vec::new();
809 let mut requested_after_timestamps = Vec::new();
810 let mut settled_after_timestamps = Vec::new();
811 let mut rated_after_inputs: Vec<(String, Option<telemetry_events::EditPredictionRating>)> =
812 Vec::new();
813 let mut accepted_after_timestamps = Vec::new();
814 let mut file_inputs = Vec::new();
815
816 for input in &args.inputs {
817 let input_string = input.to_string_lossy();
818 if let Some(timestamp) = pull_examples::parse_captured_after_input(input_string.as_ref()) {
819 captured_after_timestamps.push(timestamp.to_string());
820 } else if let Some((explicit, timestamp)) =
821 pull_examples::parse_rejected_after_input(input_string.as_ref())
822 {
823 rejected_after_timestamps.push((explicit, timestamp.to_string()));
824 } else if let Some(timestamp) =
825 pull_examples::parse_accepted_after_input(input_string.as_ref())
826 {
827 accepted_after_timestamps.push(timestamp.to_string());
828 } else if let Some(timestamp) =
829 pull_examples::parse_requested_after_input(input_string.as_ref())
830 {
831 requested_after_timestamps.push(timestamp.to_string());
832 } else if let Some(timestamp) = parse_settled_after_input(input_string.as_ref()) {
833 settled_after_timestamps.push(timestamp.to_string());
834 } else if let Some((timestamp, rating_filter)) =
835 pull_examples::parse_rated_after_input(input_string.as_ref())
836 {
837 rated_after_inputs.push((timestamp.to_string(), rating_filter));
838 } else {
839 file_inputs.push(input.clone());
840 }
841 }
842
843 let mut examples = read_example_files(&file_inputs);
844
845 // Apply offset to file examples first, then pass remaining offset to Snowflake.
846 let file_example_count = examples.len();
847 let remaining_offset = if let Some(offset) = args.offset {
848 if offset >= file_example_count {
849 examples.clear();
850 offset - file_example_count
851 } else {
852 examples.splice(0..offset, []);
853 0
854 }
855 } else {
856 0
857 };
858
859 Progress::global().set_total_examples(examples.len());
860
861 let remaining_limit_for_snowflake =
862 args.limit.map(|limit| limit.saturating_sub(examples.len()));
863
864 if let Some(0) = remaining_limit_for_snowflake {
865 log::info!(
866 "skipping Snowflake inputs because --limit is already satisfied by example files"
867 );
868 } else {
869 let max_rows_per_timestamp = remaining_limit_for_snowflake;
870
871 if !rejected_after_timestamps.is_empty() {
872 rejected_after_timestamps.sort();
873
874 let mut rejected_examples = pull_examples::fetch_rejected_examples_after(
875 http_client.clone(),
876 &rejected_after_timestamps,
877 max_rows_per_timestamp,
878 remaining_offset,
879 background_executor.clone(),
880 Some(MIN_CAPTURE_VERSION),
881 )
882 .await?;
883 examples.append(&mut rejected_examples);
884 }
885
886 if !accepted_after_timestamps.is_empty() {
887 accepted_after_timestamps.sort();
888
889 let mut accepted_examples = pull_examples::fetch_accepted_examples_after(
890 http_client.clone(),
891 &accepted_after_timestamps,
892 max_rows_per_timestamp,
893 remaining_offset,
894 background_executor.clone(),
895 Some(MIN_CAPTURE_VERSION),
896 )
897 .await?;
898 examples.append(&mut accepted_examples);
899 }
900
901 if !requested_after_timestamps.is_empty() {
902 requested_after_timestamps.sort();
903
904 let mut requested_examples = pull_examples::fetch_requested_examples_after(
905 http_client.clone(),
906 &requested_after_timestamps,
907 max_rows_per_timestamp,
908 remaining_offset,
909 background_executor.clone(),
910 Some(MIN_CAPTURE_VERSION),
911 )
912 .await?;
913 examples.append(&mut requested_examples);
914 }
915
916 if !captured_after_timestamps.is_empty() {
917 captured_after_timestamps.sort();
918
919 let mut captured_examples = pull_examples::fetch_captured_examples_after(
920 http_client.clone(),
921 &captured_after_timestamps,
922 max_rows_per_timestamp,
923 remaining_offset,
924 background_executor.clone(),
925 Some(MIN_CAPTURE_VERSION),
926 )
927 .await?;
928 examples.append(&mut captured_examples);
929 }
930
931 if !settled_after_timestamps.is_empty() {
932 settled_after_timestamps.sort();
933
934 let mut settled_examples = fetch_settled_examples_after(
935 http_client.clone(),
936 &settled_after_timestamps,
937 max_rows_per_timestamp,
938 remaining_offset,
939 background_executor.clone(),
940 Some(MIN_CAPTURE_VERSION),
941 )
942 .await?;
943 examples.append(&mut settled_examples);
944 }
945
946 if !rated_after_inputs.is_empty() {
947 rated_after_inputs.sort();
948
949 let mut rated_examples = pull_examples::fetch_rated_examples_after(
950 http_client,
951 &rated_after_inputs,
952 max_rows_per_timestamp,
953 remaining_offset,
954 background_executor,
955 Some(MIN_CAPTURE_VERSION),
956 )
957 .await?;
958 examples.append(&mut rated_examples);
959 }
960 }
961
962 crate::example::sort_examples_by_repo_and_rev(&mut examples);
963
964 if let Some(name_filter) = &args.name {
965 examples.retain(|example| example.spec.name.contains(name_filter));
966 }
967 if let Some(repo_filter) = &args.repo {
968 examples.retain(|example| example.spec.repository_url.contains(repo_filter));
969 }
970
971 // Skip resume logic for --in-place since input and output are the same file,
972 // which would incorrectly treat all input examples as already processed.
973 if !args.in_place {
974 if let Some(path) = output_path
975 && let Some(command) = &args.command
976 {
977 resume_from_output(path, &mut examples, command);
978 }
979 }
980
981 if let Some(max_duplicates) = args.max_duplicates {
982 deduplicate_examples(&mut examples, max_duplicates);
983 }
984
985 if let Some(limit) = args.limit {
986 examples.truncate(limit);
987 }
988
989 let progress = Progress::global();
990 progress.set_total_examples(examples.len());
991 progress.set_max_example_name_len(examples.iter().map(|e| &e.spec.name));
992
993 Ok(examples)
994}
995
996fn spec_hash(spec: &edit_prediction::example_spec::ExampleSpec) -> u64 {
997 let mut hasher = collections::FxHasher::default();
998 spec.hash(&mut hasher);
999 hasher.finish()
1000}
1001
1002fn chunk_examples(examples: Vec<Example>, max_parallelism: usize) -> VecDeque<Vec<Example>> {
1003 if examples.is_empty() || max_parallelism == 0 {
1004 return VecDeque::new();
1005 }
1006
1007 let chunk_size = examples.len().div_ceil(max_parallelism);
1008 examples
1009 .chunks(chunk_size)
1010 .map(|chunk| chunk.to_vec())
1011 .collect()
1012}
1013
1014fn resume_from_output(path: &PathBuf, examples: &mut Vec<Example>, command: &Command) {
1015 let file = match File::open(path) {
1016 Ok(f) => f,
1017 Err(_) => return,
1018 };
1019
1020 let input_hashes: HashSet<u64> = examples.iter().map(|e| spec_hash(&e.spec)).collect();
1021
1022 let reader = BufReader::new(file);
1023 let mut kept_lines = Vec::new();
1024 let mut kept_hashes = HashSet::default();
1025
1026 for line in reader.lines() {
1027 let line = match line {
1028 Ok(l) => l,
1029 Err(_) => continue,
1030 };
1031
1032 if let Ok(output_example) = serde_json::from_str::<Example>(&line) {
1033 let hash = spec_hash(&output_example.spec);
1034 if input_hashes.contains(&hash) && !kept_hashes.contains(&hash) {
1035 let is_complete = match command {
1036 Command::Qa(_) => output_example
1037 .qa
1038 .first()
1039 .and_then(|q| q.as_ref())
1040 .and_then(|q| q.confidence)
1041 .is_some(),
1042 Command::Repair(_) => output_example.predictions.iter().any(|p| {
1043 p.provider == PredictionProvider::Repair && p.actual_patch.is_some()
1044 }),
1045 _ => true,
1046 };
1047 if is_complete {
1048 kept_hashes.insert(hash);
1049 kept_lines.push(line);
1050 }
1051 }
1052 }
1053 }
1054
1055 let total = examples.len();
1056 let already_processed = kept_hashes.len();
1057
1058 eprintln!(
1059 "Resuming: {}/{} examples already processed",
1060 already_processed, total
1061 );
1062
1063 let file = OpenOptions::new()
1064 .write(true)
1065 .truncate(true)
1066 .open(path)
1067 .expect("Failed to open output file for rewriting");
1068 let mut writer = BufWriter::new(file);
1069 for line in &kept_lines {
1070 writeln!(writer, "{}", line).expect("Failed to write to output file");
1071 }
1072 writer.flush().expect("Failed to flush output file");
1073
1074 examples.retain(|e| !kept_hashes.contains(&spec_hash(&e.spec)));
1075}
1076
1077fn main() {
1078 let args = EpArgs::parse();
1079
1080 if args.printenv {
1081 ::util::shell_env::print_env();
1082 return;
1083 }
1084
1085 let output = args.output_path();
1086
1087 if args.markdown && output.is_none() {
1088 eprintln!("--markdown requires -o to specify the output directory");
1089 std::process::exit(1);
1090 }
1091
1092 let command = match &args.command {
1093 Some(cmd) => cmd.clone(),
1094 None => {
1095 EpArgs::command().print_help().unwrap();
1096 return;
1097 }
1098 };
1099
1100 match &command {
1101 Command::ImportBatch(import_args) => {
1102 gpui::block_on(async {
1103 match import_args.provider {
1104 BatchProvider::Anthropic => {
1105 let client = anthropic_client::AnthropicClient::batch(&paths::LLM_CACHE_DB)
1106 .expect("Failed to create Anthropic client");
1107 if let Err(e) = client.import_batches(&import_args.batch_ids).await {
1108 eprintln!("Error importing Anthropic batches: {:?}", e);
1109 std::process::exit(1);
1110 }
1111 }
1112 BatchProvider::Openai => {
1113 let client = openai_client::OpenAiClient::batch(&paths::LLM_CACHE_DB)
1114 .expect("Failed to create OpenAI client");
1115 if let Err(e) = client.import_batches(&import_args.batch_ids).await {
1116 eprintln!("Error importing OpenAI batches: {:?}", e);
1117 std::process::exit(1);
1118 }
1119 }
1120 }
1121 println!(
1122 "Successfully imported {} batch(es)",
1123 import_args.batch_ids.len()
1124 );
1125 });
1126 return;
1127 }
1128 Command::Clean => {
1129 std::fs::remove_dir_all(&*paths::DATA_DIR).unwrap();
1130 return;
1131 }
1132 Command::PrintZetaFormats => {
1133 use strum::IntoEnumIterator as _;
1134 for format in ZetaFormat::iter() {
1135 println!("{}", format.to_string().to_lowercase());
1136 }
1137 return;
1138 }
1139
1140 Command::Synthesize(synth_args) => {
1141 let output_dir = if let Some(output_dir) = args.output {
1142 output_dir
1143 } else {
1144 let default_output_dir = env::current_dir()
1145 .unwrap()
1146 .join("crates/edit_prediction_cli/evals-generated");
1147 if default_output_dir.parent().unwrap().exists() {
1148 std::fs::create_dir(&default_output_dir).ok();
1149 default_output_dir
1150 } else {
1151 panic!("output dir is required");
1152 }
1153 };
1154 let config = SynthesizeConfig {
1155 repo_urls: synth_args.repos.clone(),
1156 count: synth_args.count,
1157 max_commits: synth_args.max_commits,
1158 output_dir,
1159 fresh: synth_args.fresh,
1160 };
1161 gpui::block_on(async {
1162 if let Err(e) = run_synthesize(config).await {
1163 eprintln!("Error: {:?}", e);
1164 std::process::exit(1);
1165 }
1166 });
1167 return;
1168 }
1169 Command::SplitCommit(split_commit_args) => {
1170 if let Err(error) = split_commit::run_split_commit(
1171 split_commit_args,
1172 &args.inputs,
1173 output.as_ref(),
1174 args.failed,
1175 ) {
1176 eprintln!("{error:#}");
1177 std::process::exit(1);
1178 }
1179 return;
1180 }
1181 Command::TruncatePatch(truncate_args) => {
1182 if let Err(error) =
1183 truncate_expected_patch::run_truncate_expected_patch(truncate_args, &args.inputs)
1184 {
1185 eprintln!("{error:#}");
1186 std::process::exit(1);
1187 }
1188 return;
1189 }
1190 Command::Split(split_args) => {
1191 if let Err(error) = split_dataset::run_split(split_args, &args.inputs) {
1192 eprintln!("{error:#}");
1193 std::process::exit(1);
1194 }
1195 return;
1196 }
1197 Command::FilterLanguages(filter_args) => {
1198 if let Err(error) =
1199 run_filter_languages(filter_args, &args.inputs, args.output.as_ref())
1200 {
1201 eprintln!("{error:#}");
1202 std::process::exit(1);
1203 }
1204 return;
1205 }
1206
1207 _ => {}
1208 }
1209
1210 let http_client = Arc::new(ReqwestClient::new());
1211 let app = gpui_platform::headless().with_http_client(http_client);
1212
1213 app.run(move |cx| {
1214 let app_state = Arc::new(headless::init(cx));
1215 EditPredictionStore::global(&app_state.client, &app_state.user_store, cx);
1216
1217 cx.spawn(async move |cx| {
1218 let result = async {
1219 let examples = load_examples(
1220 app_state.client.http_client(),
1221 &args,
1222 output.as_ref(),
1223 cx.background_executor().clone(),
1224 )
1225 .await?;
1226
1227 match &command {
1228 Command::Predict(args) | Command::Score(args) => {
1229 predict::sync_batches(args.provider.as_ref()).await?;
1230 }
1231 Command::Eval(args) => {
1232 if !args.context_only {
1233 predict::sync_batches(args.predict.provider.as_ref()).await?;
1234 }
1235 }
1236 Command::Qa(args) => {
1237 qa::sync_batches(args).await?;
1238 }
1239 Command::Repair(args) => {
1240 repair::sync_batches(args).await?;
1241 }
1242 _ => (),
1243 }
1244
1245 let failfast_on_single_example = examples.len() == 1;
1246
1247 // For --markdown mode, create the output directory if it doesn't exist
1248 if args.markdown {
1249 let dir = output.as_ref().expect("--markdown requires -o");
1250 if !dir.exists() {
1251 std::fs::create_dir_all(dir)
1252 .expect("Failed to create markdown output directory");
1253 }
1254 }
1255
1256 // Set up JSONL output writer (not used in markdown mode)
1257 let mut output_sender: Option<mpsc::UnboundedSender<String>> = None;
1258 let mut in_place_temp_path: Option<PathBuf> = None;
1259 if !args.markdown
1260 && let Some(output_path) = output.as_ref()
1261 {
1262 let write_path = if args.in_place {
1263 let temp = output_path.with_extension("jsonl.tmp");
1264 in_place_temp_path = Some(temp.clone());
1265 temp
1266 } else {
1267 output_path.clone()
1268 };
1269
1270 let file = OpenOptions::new()
1271 .create(true)
1272 .write(true)
1273 .truncate(args.in_place)
1274 .append(!args.in_place)
1275 .open(&write_path)
1276 .expect("Failed to open output file");
1277
1278 let mut writer = BufWriter::new(file);
1279 let (sender, mut receiver) = mpsc::unbounded::<String>();
1280 cx.background_spawn(async move {
1281 while let Some(line) = receiver.next().await {
1282 writeln!(writer, "{}", line).expect("Failed to write example");
1283 writer.flush().expect("Failed to flush output");
1284 }
1285 })
1286 .detach();
1287 output_sender = Some(sender);
1288 }
1289
1290 let example_batches = if args.group_by_repo {
1291 group_examples_by_repo(examples)
1292 } else {
1293 chunk_examples(examples, args.max_parallelism)
1294 };
1295 let example_batches = Mutex::new(example_batches);
1296 let finished_examples = Mutex::new(Vec::new());
1297
1298 let mut tasks = Vec::new();
1299 for _ in 0..args.max_parallelism {
1300 tasks.push(async {
1301 loop {
1302 let Some(mut repo_examples) =
1303 example_batches.lock().unwrap().pop_front()
1304 else {
1305 break;
1306 };
1307 for example in &mut repo_examples {
1308 let example_progress =
1309 Progress::global().start_group(&example.spec.name);
1310
1311 let result = async {
1312 match &command {
1313 Command::Read(_) => {}
1314 Command::LoadProject => {
1315 run_load_project(
1316 example,
1317 app_state.clone(),
1318 &example_progress,
1319 cx.clone(),
1320 )
1321 .await?;
1322 }
1323 Command::Context(args) => {
1324 run_context_retrieval(
1325 example,
1326 app_state.clone(),
1327 &example_progress,
1328 args.context_types(),
1329 args.force,
1330 cx.clone(),
1331 )
1332 .await?;
1333 }
1334 Command::FormatPrompt(args) => {
1335 run_format_prompt(
1336 example,
1337 args,
1338 app_state.clone(),
1339 &example_progress,
1340 cx.clone(),
1341 )
1342 .await?;
1343 }
1344 Command::Predict(args) => {
1345 run_prediction(
1346 example,
1347 args,
1348 app_state.clone(),
1349 &example_progress,
1350 cx.clone(),
1351 )
1352 .await?;
1353 }
1354 Command::ParseOutput => {
1355 parse_output::run_parse_output(example)?;
1356 }
1357 Command::Distill => {
1358 run_distill(example).await?;
1359 }
1360 Command::Score(args) => {
1361 run_scoring(
1362 example,
1363 args,
1364 app_state.clone(),
1365 &example_progress,
1366 cx.clone(),
1367 false,
1368 None,
1369 None,
1370 )
1371 .await?;
1372 }
1373 Command::Eval(args) => {
1374 let context_source_filter =
1375 args.context_source_filter();
1376 if args.context_only {
1377 score::run_context_coverage_scoring(
1378 example,
1379 &example_progress,
1380 Some(args.related_context_limit * 3),
1381 context_source_filter.as_deref(),
1382 )?;
1383 } else {
1384 run_scoring(
1385 example,
1386 &args.predict,
1387 app_state.clone(),
1388 &example_progress,
1389 cx.clone(),
1390 true,
1391 Some(args.related_context_limit * 3),
1392 context_source_filter,
1393 )
1394 .await?;
1395 }
1396 }
1397 Command::Qa(args) => {
1398 qa::run_qa(example, args, &example_progress).await?;
1399 }
1400 Command::Repair(args) => {
1401 repair::run_repair(example, args, &example_progress)
1402 .await?;
1403 }
1404 Command::Clean
1405 | Command::Synthesize(_)
1406 | Command::SplitCommit(_)
1407 | Command::Split(_)
1408 | Command::TruncatePatch(_)
1409 | Command::FilterLanguages(_)
1410 | Command::ImportBatch(_)
1411 | Command::PrintZetaFormats => {
1412 unreachable!()
1413 }
1414 }
1415 anyhow::Ok(())
1416 }
1417 .await;
1418
1419 let failed = if let Err(error) = result {
1420 handle_error(
1421 error,
1422 &args,
1423 &command,
1424 &app_state,
1425 failfast_on_single_example,
1426 &example,
1427 )
1428 .await;
1429 true
1430 } else {
1431 false
1432 };
1433
1434 let should_write = !failed || args.failed == FailedHandling::Keep;
1435 if should_write {
1436 if args.markdown {
1437 let markdown_dir =
1438 output.as_ref().expect("--markdown requires -o");
1439 let filename = format!("{}.md", example.spec.filename());
1440 let path = markdown_dir.join(&filename);
1441 let markdown = example.spec.to_markdown();
1442 std::fs::write(&path, &markdown)
1443 .expect("Failed to write markdown file");
1444 } else if let Some(ref mut sender) = output_sender.clone() {
1445 let line = serde_json::to_string(&example).unwrap();
1446 sender
1447 .send(line)
1448 .await
1449 .expect("Failed to send to output writer");
1450 } else if args.output.is_none()
1451 && !matches!(command, Command::Eval(_))
1452 {
1453 let line = serde_json::to_string(&example).unwrap();
1454 println!("{}", line);
1455 }
1456 }
1457 }
1458
1459 let project = repo_examples
1460 .iter()
1461 .find_map(|e| e.state.as_ref().map(|s| s.project.clone()));
1462
1463 if let Some(project) = project {
1464 let mut cx = cx.clone();
1465
1466 let shutdown_task: Task<()> =
1467 project.update(&mut cx, |project, cx| {
1468 let lsp_store = project.lsp_store();
1469 lsp_store.update(cx, |lsp_store, cx| {
1470 lsp_store.shutdown_all_language_servers(cx)
1471 })
1472 });
1473
1474 shutdown_task.await;
1475
1476 if let Some(ep_store) =
1477 cx.update(|cx| EditPredictionStore::try_global(cx))
1478 {
1479 ep_store.update(&mut cx, |store, _| {
1480 store.remove_project(&project);
1481 });
1482 }
1483 }
1484
1485 for example in &mut repo_examples {
1486 example.state.take();
1487 }
1488 finished_examples
1489 .lock()
1490 .unwrap()
1491 .extend_from_slice(&repo_examples);
1492 }
1493 });
1494 }
1495 futures::future::join_all(tasks).await;
1496
1497 Progress::global().finalize();
1498
1499 let is_markdown = args.markdown;
1500 let write_path = in_place_temp_path.as_ref().or(output.as_ref());
1501 match &command {
1502 Command::Predict(args) | Command::Score(args) => {
1503 predict::sync_batches(args.provider.as_ref()).await?;
1504 if args.wait {
1505 predict::wait_for_batches(args.provider.as_ref()).await?;
1506 let mut examples =
1507 std::mem::take(&mut *finished_examples.lock().unwrap());
1508 predict::reprocess_after_batch_wait(&mut examples, args).await?;
1509 rewrite_output(&examples, write_path, is_markdown)?;
1510 *finished_examples.lock().unwrap() = examples;
1511 }
1512 }
1513 Command::Eval(args) => {
1514 if !args.context_only {
1515 predict::sync_batches(args.predict.provider.as_ref()).await?;
1516 if args.predict.wait {
1517 predict::wait_for_batches(args.predict.provider.as_ref()).await?;
1518 let mut examples =
1519 std::mem::take(&mut *finished_examples.lock().unwrap());
1520 predict::reprocess_after_batch_wait(&mut examples, &args.predict)
1521 .await?;
1522 rewrite_output(&examples, write_path, is_markdown)?;
1523 *finished_examples.lock().unwrap() = examples;
1524 }
1525 }
1526 }
1527 Command::Qa(args) => {
1528 qa::sync_batches(args).await?;
1529 }
1530 Command::Repair(args) => {
1531 repair::sync_batches(args).await?;
1532 if args.wait {
1533 repair::wait_for_batches(args).await?;
1534 let mut examples =
1535 std::mem::take(&mut *finished_examples.lock().unwrap());
1536 repair::reprocess_after_batch_wait(&mut examples, args).await?;
1537 rewrite_output(&examples, write_path, is_markdown)?;
1538 *finished_examples.lock().unwrap() = examples;
1539 }
1540 }
1541 _ => (),
1542 }
1543
1544 match &command {
1545 Command::Eval(args) => {
1546 let examples = finished_examples.lock().unwrap();
1547 let context_source_filter = args.context_source_filter();
1548 score::print_report(
1549 &examples,
1550 args.verbose,
1551 args.context_only,
1552 Some(args.related_context_limit * 3),
1553 context_source_filter.as_deref(),
1554 );
1555 if let Some(summary_path) = &args.summary_json {
1556 score::write_summary_json(
1557 &examples,
1558 summary_path,
1559 Some(args.related_context_limit * 3),
1560 context_source_filter.as_deref(),
1561 )?;
1562 }
1563 }
1564 Command::Repair(args) => {
1565 let examples = finished_examples.lock().unwrap();
1566 repair::print_report(&examples, args.confidence_threshold);
1567 }
1568 _ => (),
1569 };
1570
1571 // For --in-place, atomically rename temp file to original
1572 if let Some(temp_path) = &in_place_temp_path {
1573 let final_path = output.as_ref().expect("in_place_temp_path requires output");
1574 std::fs::rename(temp_path, final_path)
1575 .expect("Failed to rename temp file to final output");
1576 }
1577
1578 anyhow::Ok(())
1579 }
1580 .await;
1581
1582 if let Err(e) = result {
1583 panic!("Fatal error: {:?}", e);
1584 }
1585
1586 let _ = cx.update(|cx| cx.quit());
1587 })
1588 .detach();
1589 });
1590}
1591
1592fn rewrite_output(
1593 examples: &[Example],
1594 output_path: Option<&PathBuf>,
1595 markdown: bool,
1596) -> anyhow::Result<()> {
1597 if markdown {
1598 let dir = output_path.context("--markdown requires -o")?;
1599 for example in examples {
1600 let filename = format!("{}.md", example.spec.filename());
1601 let path = dir.join(&filename);
1602 let markdown = example.spec.to_markdown();
1603 std::fs::write(&path, &markdown).context("Failed to write markdown file")?;
1604 }
1605 } else if let Some(path) = output_path {
1606 let file = OpenOptions::new()
1607 .create(true)
1608 .write(true)
1609 .truncate(true)
1610 .open(path)
1611 .context("Failed to open output file for rewriting")?;
1612 let mut writer = BufWriter::new(file);
1613 for example in examples {
1614 let line = serde_json::to_string(example)?;
1615 writeln!(writer, "{}", line)?;
1616 }
1617 writer.flush()?;
1618 } else {
1619 for example in examples {
1620 let line = serde_json::to_string(example)?;
1621 println!("{}", line);
1622 }
1623 }
1624 Ok(())
1625}
1626
1627async fn handle_error(
1628 error: anyhow::Error,
1629 args: &EpArgs,
1630 command: &Command,
1631 app_state: &Arc<headless::EpAppState>,
1632 failfast_on_single_example: bool,
1633 example: &Example,
1634) {
1635 Progress::global().increment_failed();
1636
1637 let msg;
1638 if !matches!(args.failed, FailedHandling::SkipNoFiles) {
1639 let example_name = example.spec.filename();
1640
1641 let failed_example_path = FAILED_EXAMPLES_DIR.join(format!("{}.json", example_name));
1642 app_state
1643 .fs
1644 .write(
1645 &failed_example_path,
1646 &serde_json::to_vec_pretty(&example).unwrap(),
1647 )
1648 .await
1649 .unwrap();
1650 let err_path = FAILED_EXAMPLES_DIR.join(format!("{}_err.txt", example_name));
1651 app_state
1652 .fs
1653 .write(&err_path, format!("{error:?}").as_bytes())
1654 .await
1655 .unwrap();
1656
1657 let failed_jsonl_path = RUN_DIR.join("failed.jsonl");
1658 let mut file = OpenOptions::new()
1659 .create(true)
1660 .append(true)
1661 .open(&failed_jsonl_path)
1662 .expect("Failed to open failed.jsonl");
1663 writeln!(file, "{}", serde_json::to_string(example).unwrap())
1664 .expect("Failed to write to failed.jsonl");
1665
1666 let cursor_path = match example.repo_name() {
1667 Ok(repo_name) => repo_name.worktree_path().join(&example.spec.cursor_path),
1668 Err(_) => example.spec.cursor_path.as_ref().to_path_buf(),
1669 };
1670 msg = format!(
1671 indoc::indoc! {"
1672 While processing \"{}\":
1673
1674 \x1b[31m{:?}\x1b[0m
1675
1676 Example: \x1b[36m{}\x1b[0m
1677 Error file: \x1b[36m{}\x1b[0m
1678 Cursor file: \x1b[36m{}\x1b[0m
1679 Re-run: cargo run -p edit_prediction_cli -- {} \x1b[36m{}\x1b[0m
1680 "},
1681 example.spec.name,
1682 error,
1683 failed_example_path.display(),
1684 err_path.display(),
1685 cursor_path.display(),
1686 command,
1687 failed_example_path.display(),
1688 );
1689 } else {
1690 msg = format!(
1691 indoc::indoc! {"
1692 While processing \"{}\":
1693
1694 \x1b[31m{:?}\x1b[0m
1695 "},
1696 example.spec.name, error
1697 );
1698 }
1699
1700 if args.failfast || failfast_on_single_example {
1701 Progress::global().finalize();
1702 panic!("{}", msg);
1703 } else {
1704 log::error!("{}", msg);
1705 }
1706}
1707