1use serde::{Deserialize, Serialize};
7use std::collections::HashMap;
8
9#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
11pub struct ModelEntry {
12 pub provider: String,
14
15 pub model: String,
17}
18
19impl ModelEntry {
20 pub fn new(provider: String, model: String) -> Self {
22 Self { provider, model }
23 }
24}
25
26#[derive(Debug, Clone, Serialize, Deserialize)]
34pub struct ModelConfig {
35 #[serde(default)]
37 pub models: Vec<ModelEntry>,
38
39 #[serde(default = "default_allow_user_default")]
42 pub allow_user_default: bool,
43
44 #[serde(default)]
46 pub parameters: HashMap<String, serde_json::Value>,
47
48 #[serde(default)]
56 pub request_timeout_secs: Option<u64>,
57}
58
59fn default_allow_user_default() -> bool {
60 true
61}
62
63#[derive(Debug, Clone, PartialEq)]
81pub enum OutputCap {
82 Tokens(usize),
84 WindowPercent(f64),
86 RegionPercent {
88 percent: f64,
90 region: String,
92 },
93}
94
95impl OutputCap {
96 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 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 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 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 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 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 pub fn provider(&self) -> &str {
213 self.models
214 .first()
215 .map(|e| e.provider.as_str())
216 .unwrap_or("anthropic")
217 }
218
219 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 #[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}