Skip to main content

leviath_core/blueprint/
model.rs

1//! Which model a stage runs on, and what a user may override.
2//!
3//! Two levels: a [`ModelEntry`] names a provider and model, and [`ModelConfig`]
4//! decides whether the user's own default may stand in for it.
5
6use serde::{Deserialize, Serialize};
7use std::collections::HashMap;
8
9/// A single model entry within a [`ModelConfig`] models list.
10#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
11pub struct ModelEntry {
12    /// Provider name (e.g., "anthropic", "openai")
13    pub provider: String,
14
15    /// Model identifier (e.g., "claude-sonnet-4-6")
16    pub model: String,
17}
18
19impl ModelEntry {
20    /// One provider/model pair in a stage's fallback list.
21    pub fn new(provider: String, model: String) -> Self {
22        Self { provider, model }
23    }
24}
25
26/// Model configuration for a stage.
27///
28/// Models are specified as an ordered priority list in `models`. The first
29/// entry whose provider is registered at runtime is used. When
30/// `allow_user_default` is true (the default), the user's configured default
31/// model is tried as a last resort. When false, the stage fails if none of
32/// the listed models are available.
33#[derive(Debug, Clone, Serialize, Deserialize)]
34pub struct ModelConfig {
35    /// Ordered list of models to try (first available wins).
36    #[serde(default)]
37    pub models: Vec<ModelEntry>,
38
39    /// When true (default), fall back to the user's configured default model
40    /// if none of the listed models are available.
41    #[serde(default = "default_allow_user_default")]
42    pub allow_user_default: bool,
43
44    /// Optional parameters that apply to whichever model gets selected.
45    #[serde(default)]
46    pub parameters: HashMap<String, serde_json::Value>,
47
48    /// Optional per-stage cap on the wall-clock time (in seconds) one inference
49    /// for this stage may run - the whole call including retries. When set, it
50    /// overrides the default job timeout; when `None`, the default applies.
51    ///
52    /// This lets a stage with slow first-token latency (e.g. a large-prompt
53    /// analyze call) get a long cap while a quick iterative stage fails fast on
54    /// a stalled connection instead of hanging for the full default.
55    #[serde(default)]
56    pub request_timeout_secs: Option<u64>,
57}
58
59fn default_allow_user_default() -> bool {
60    true
61}
62
63/// How large one reply may be, as a stage's `parameters.max_output_tokens`
64/// says it.
65///
66/// A bare number is the classic form and is sent as written. The other two
67/// are relative, because a fixed number is the wrong shape for the question:
68/// a stage that rewrites a whole report needs "as much as this model can
69/// give", and a stage that fills a region needs "as much as the region
70/// holds". A fixed `24000` on the report stage was smaller than the report,
71/// and every reply was cut off. A relative cap is clamped to the model's own
72/// maximum, since asking for more than that is refused by the provider.
73///
74/// ```toml
75/// parameters = { max_output_tokens = 8000 }              # tokens
76/// parameters = { max_output_tokens = "40%" }             # of the model's context window
77/// parameters = { max_output_tokens = "100% of claims" }  # of the `claims` region's budget
78/// parameters = { max_output_tokens = { percent = 100, of = "claims" } }
79/// ```
80#[derive(Debug, Clone, PartialEq)]
81pub enum OutputCap {
82    /// A fixed number of output tokens, sent as written.
83    Tokens(usize),
84    /// A fraction (`0.0..=1.0`) of the model's context window.
85    WindowPercent(f64),
86    /// A fraction (`0.0..=1.0`) of a named region's token budget.
87    RegionPercent {
88        /// The fraction.
89        percent: f64,
90        /// The region whose budget it is a fraction of.
91        region: String,
92    },
93}
94
95impl OutputCap {
96    /// Read the cap from the parameter's JSON value. `Err` carries the reason
97    /// in the words a blueprint author needs, for the manifest loader to
98    /// surface: a cap that does not parse must fail the load, because the
99    /// alternative (no cap) is exactly the silent no-op that hides a typo.
100    pub fn parse(value: &serde_json::Value) -> Result<Self, String> {
101        match value {
102            serde_json::Value::Number(n) => match n.as_u64() {
103                Some(t) if t > 0 => Ok(OutputCap::Tokens(t as usize)),
104                _ => Err(format!(
105                    "max_output_tokens = {n} must be a positive whole number of tokens"
106                )),
107            },
108            serde_json::Value::String(s) => Self::parse_text(s),
109            serde_json::Value::Object(table) => {
110                let percent = match table.get("percent") {
111                    Some(serde_json::Value::Number(n)) => Self::fraction(&format!("{n}%"))?,
112                    Some(serde_json::Value::String(s)) => Self::fraction(s)?,
113                    _ => {
114                        return Err(
115                            "max_output_tokens as a table needs `percent` (a number, or \
116                                    a string like \"40%\") and optionally `of = \"<region>\"`"
117                                .to_string(),
118                        );
119                    }
120                };
121                match table.get("of") {
122                    None => Ok(OutputCap::WindowPercent(percent)),
123                    Some(serde_json::Value::String(region)) if !region.trim().is_empty() => {
124                        Ok(OutputCap::RegionPercent {
125                            percent,
126                            region: region.trim().to_string(),
127                        })
128                    }
129                    Some(_) => Err("max_output_tokens: `of` must name a region".to_string()),
130                }
131            }
132            other => Err(format!(
133                "max_output_tokens = {other} is not a token count, a percentage like \"40%\", \
134                 or \"<percent>% of <region>\""
135            )),
136        }
137    }
138
139    /// `"40%"` or `"100% of claims"`.
140    fn parse_text(s: &str) -> Result<Self, String> {
141        match s.split_once(" of ") {
142            None => Ok(OutputCap::WindowPercent(Self::fraction(s)?)),
143            Some((pct, region)) => {
144                let region = region.trim();
145                if region.is_empty() {
146                    return Err(format!(
147                        "max_output_tokens = \"{s}\" names no region after `of`"
148                    ));
149                }
150                Ok(OutputCap::RegionPercent {
151                    percent: Self::fraction(pct)?,
152                    region: region.to_string(),
153                })
154            }
155        }
156    }
157
158    /// The `(0, 100]` percent rule shared with region budgets, in this
159    /// setting's words.
160    fn fraction(s: &str) -> Result<f64, String> {
161        crate::layout::BudgetSpec::parse_budget(s).map_err(|e| format!("max_output_tokens: {e}"))
162    }
163
164    /// The cap in tokens for one request.
165    ///
166    /// `model_window` and `model_max_output` are the model's own limits;
167    /// `region_budget` answers "how many tokens may region X hold" for the
168    /// window the request is built from. A relative cap is clamped to the
169    /// model's maximum. A region cap naming a region the stage does not carry
170    /// falls back to the model's maximum, which is the same "as much as you
171    /// can" the author was reaching for; the loader already warned about the
172    /// name.
173    pub fn resolve(
174        &self,
175        model_window: usize,
176        model_max_output: usize,
177        region_budget: impl Fn(&str) -> Option<usize>,
178    ) -> usize {
179        let share = |whole: usize, fraction: f64| (whole as f64 * fraction).round() as usize;
180        match self {
181            OutputCap::Tokens(t) => *t,
182            OutputCap::WindowPercent(p) => share(model_window, *p).min(model_max_output),
183            OutputCap::RegionPercent { percent, region } => match region_budget(region) {
184                Some(budget) => share(budget, *percent).min(model_max_output),
185                None => model_max_output,
186            },
187        }
188        .max(1)
189    }
190}
191
192impl ModelConfig {
193    /// The stage's output cap, if `parameters.max_output_tokens` sets one.
194    pub fn output_cap(&self) -> Result<Option<OutputCap>, String> {
195        self.parameters
196            .get("max_output_tokens")
197            .map(OutputCap::parse)
198            .transpose()
199    }
200
201    /// Create a new model configuration with a single model entry.
202    pub fn new(provider: String, model: String) -> Self {
203        Self {
204            models: vec![ModelEntry::new(provider, model)],
205            allow_user_default: true,
206            parameters: HashMap::new(),
207            request_timeout_secs: None,
208        }
209    }
210
211    /// Convenience: provider of the first model entry (for backward compat).
212    pub fn provider(&self) -> &str {
213        self.models
214            .first()
215            .map(|e| e.provider.as_str())
216            .unwrap_or("anthropic")
217    }
218
219    /// Convenience: model name of the first model entry (for backward compat).
220    pub fn model(&self) -> &str {
221        self.models
222            .first()
223            .map(|e| e.model.as_str())
224            .unwrap_or("claude-sonnet-4-6")
225    }
226}
227
228#[cfg(test)]
229mod tests {
230    use super::*;
231    use serde_json::json;
232
233    #[test]
234    fn every_written_form_of_the_cap_parses() {
235        assert_eq!(OutputCap::parse(&json!(8000)), Ok(OutputCap::Tokens(8000)));
236        assert_eq!(
237            OutputCap::parse(&json!("40%")),
238            Ok(OutputCap::WindowPercent(0.4))
239        );
240        assert_eq!(
241            OutputCap::parse(&json!("100% of claims")),
242            Ok(OutputCap::RegionPercent {
243                percent: 1.0,
244                region: "claims".to_string()
245            })
246        );
247        assert_eq!(
248            OutputCap::parse(&json!({"percent": 25})),
249            Ok(OutputCap::WindowPercent(0.25))
250        );
251        assert_eq!(
252            OutputCap::parse(&json!({"percent": "50%", "of": " report "})),
253            Ok(OutputCap::RegionPercent {
254                percent: 0.5,
255                region: "report".to_string()
256            })
257        );
258    }
259
260    /// Each rejection names what was wrong in the author's terms; none of
261    /// them is a silent "no cap".
262    #[test]
263    fn a_cap_that_does_not_parse_says_why() {
264        let err = |v: serde_json::Value| OutputCap::parse(&v).expect_err("rejected");
265        assert!(err(json!(0)).contains("positive whole number"));
266        assert!(err(json!(-5)).contains("positive whole number"));
267        assert!(err(json!("forty")).contains("must end with '%'"));
268        assert!(err(json!("150%")).contains("at most 100%"));
269        assert!(err(json!("50% of ")).contains("names no region"));
270        assert!(err(json!("x% of claims")).contains("not a valid number"));
271        assert!(err(json!({"of": "claims"})).contains("needs `percent`"));
272        assert!(err(json!({"percent": 0})).contains("greater than 0%"));
273        assert!(err(json!({"percent": "abc"})).contains("must end with"));
274        assert!(err(json!({"percent": 10, "of": 3})).contains("must name a region"));
275        assert!(err(json!({"percent": 10, "of": ""})).contains("must name a region"));
276        assert!(err(json!(true)).contains("not a token count"));
277    }
278
279    #[test]
280    fn a_cap_resolves_against_the_model_and_the_region_and_never_below_one() {
281        let budget = |name: &str| (name == "claims").then_some(3_000);
282        assert_eq!(
283            OutputCap::Tokens(70_000).resolve(200_000, 65_535, budget),
284            70_000
285        );
286        assert_eq!(
287            OutputCap::WindowPercent(0.4).resolve(200_000, 65_535, budget),
288            65_535
289        );
290        assert_eq!(
291            OutputCap::WindowPercent(0.1).resolve(200_000, 65_535, budget),
292            20_000
293        );
294        let claims = OutputCap::RegionPercent {
295            percent: 1.0,
296            region: "claims".to_string(),
297        };
298        assert_eq!(claims.resolve(200_000, 65_535, budget), 3_000);
299        assert_eq!(claims.resolve(200_000, 2_000, budget), 2_000);
300        let missing = OutputCap::RegionPercent {
301            percent: 1.0,
302            region: "gone".to_string(),
303        };
304        assert_eq!(missing.resolve(200_000, 65_535, budget), 65_535);
305        assert_eq!(OutputCap::WindowPercent(0.001).resolve(10, 10, budget), 1);
306    }
307
308    #[test]
309    fn a_model_config_reads_its_own_cap() {
310        let mut config = ModelConfig::new("p".to_string(), "m".to_string());
311        assert_eq!(config.output_cap(), Ok(None));
312        config
313            .parameters
314            .insert("max_output_tokens".to_string(), json!("30%"));
315        assert_eq!(config.output_cap(), Ok(Some(OutputCap::WindowPercent(0.3))));
316        config
317            .parameters
318            .insert("max_output_tokens".to_string(), json!("lots"));
319        assert!(config.output_cap().is_err());
320    }
321}