Skip to repository content

tenant.openagents/omega

No repository description is available.

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

main.rs

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