1use serde::{Deserialize, Serialize};
49
50#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
54#[serde(rename_all = "snake_case")]
55#[non_exhaustive]
56pub enum LimitSource {
57 User,
62 Endpoint,
64 Preset,
66}
67
68impl LimitSource {
69 pub const fn as_str(self) -> &'static str {
71 match self {
72 LimitSource::User => "user",
73 LimitSource::Endpoint => "endpoint",
74 LimitSource::Preset => "preset",
75 }
76 }
77
78 pub fn parse(s: &str) -> Option<Self> {
80 match s.trim() {
81 "user" => Some(LimitSource::User),
82 "endpoint" => Some(LimitSource::Endpoint),
83 "preset" => Some(LimitSource::Preset),
84 _ => None,
85 }
86 }
87}
88
89#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
95#[serde(rename_all = "camelCase")]
96#[non_exhaustive]
97pub struct TokenLimits {
98 pub context_window: Option<u32>,
102 pub max_output: Option<u32>,
107 pub source: LimitSource,
112 pub context_window_source: Option<LimitSource>,
114 pub max_output_source: Option<LimitSource>,
116}
117
118const fn source_of(value: Option<u32>, source: LimitSource) -> Option<LimitSource> {
120 match value {
121 Some(_) => Some(source),
122 None => None,
123 }
124}
125
126impl TokenLimits {
127 pub const fn from_endpoint(context_window: Option<u32>, max_output: Option<u32>) -> Self {
129 Self::with_source(context_window, max_output, LimitSource::Endpoint)
130 }
131
132 pub const fn from_preset(context_window: Option<u32>, max_output: Option<u32>) -> Self {
134 Self::with_source(context_window, max_output, LimitSource::Preset)
135 }
136
137 pub const fn from_user(context_window: Option<u32>, max_output: Option<u32>) -> Self {
139 Self::with_source(context_window, max_output, LimitSource::User)
140 }
141
142 pub const fn with_source(
146 context_window: Option<u32>,
147 max_output: Option<u32>,
148 source: LimitSource,
149 ) -> Self {
150 Self {
151 context_window,
152 max_output,
153 source,
154 context_window_source: source_of(context_window, source),
155 max_output_source: source_of(max_output, source),
156 }
157 }
158
159 pub const fn or(self, fallback: TokenLimits) -> TokenLimits {
178 if self.is_empty() {
179 return fallback;
180 }
181 let (context_window, context_window_source) = match self.context_window {
182 Some(v) => (Some(v), self.context_window_source),
183 None => (fallback.context_window, fallback.context_window_source),
184 };
185 let (max_output, max_output_source) = match self.max_output {
186 Some(v) => (Some(v), self.max_output_source),
187 None => (fallback.max_output, fallback.max_output_source),
188 };
189 TokenLimits {
190 context_window,
191 max_output,
192 source: self.source,
193 context_window_source,
194 max_output_source,
195 }
196 }
197
198 pub const fn is_empty(&self) -> bool {
200 self.context_window.is_none() && self.max_output.is_none()
201 }
202
203 pub fn input_budget(&self, reserve_output: u32) -> Option<u32> {
223 self.context_window
224 .map(|w| w.saturating_sub(reserve_output))
225 }
226}
227
228pub const CONTEXT_WINDOW_FIELDS: &[&str] = &[
233 "context_length",
234 "context_window",
235 "max_context_length",
236 "max_input_tokens",
237];
238
239pub const MAX_OUTPUT_FIELDS: &[&str] =
241 &["max_completion_tokens", "max_output_tokens", "max_tokens"];
242
243pub fn parse_model_limits(item: &serde_json::Value) -> Option<TokenLimits> {
262 fn num(v: Option<&serde_json::Value>) -> Option<u32> {
264 let v = v?;
265 let n = v
266 .as_u64()
267 .or_else(|| v.as_f64().filter(|f| *f >= 0.0).map(|f| f as u64))
268 .or_else(|| v.as_str()?.trim().parse::<u64>().ok())?;
269 if n == 0 {
271 return None;
272 }
273 u32::try_from(n).ok()
274 }
275
276 let top = item.get("top_provider");
277 let pick = |keys: &[&str]| -> Option<u32> {
278 for k in keys {
280 if let Some(n) = num(top.and_then(|t| t.get(*k))) {
281 return Some(n);
282 }
283 }
284 keys.iter().find_map(|k| num(item.get(*k)))
285 };
286
287 let context_window = pick(CONTEXT_WINDOW_FIELDS);
288 let max_output = pick(MAX_OUTPUT_FIELDS);
289
290 if context_window.is_none() && max_output.is_none() {
291 return None;
292 }
293 Some(TokenLimits::from_endpoint(context_window, max_output))
294}
295
296#[cfg(test)]
297mod tests {
298 use super::*;
299 use serde_json::json;
300
301 #[test]
307 fn nested_top_provider_wins() {
308 let item = json!({
309 "id": "some/model",
310 "context_length": 1_048_576,
311 "top_provider": { "context_length": 131_072, "max_completion_tokens": 32_768 }
312 });
313 let l = parse_model_limits(&item).unwrap();
314 assert_eq!(l.context_window, Some(131_072), "该取服务商实际提供的");
315 assert_eq!(l.max_output, Some(32_768));
316 assert_eq!(l.source, LimitSource::Endpoint);
317 }
318
319 #[test]
321 fn falls_back_to_top_level() {
322 let item = json!({ "id": "m", "context_length": 200_000 });
323 let l = parse_model_limits(&item).unwrap();
324 assert_eq!(l.context_window, Some(200_000));
325 assert_eq!(l.max_output, None, "没报就是没报,不猜");
326 }
327
328 #[test]
333 fn bare_openai_shape_yields_none() {
334 let item = json!({ "id": "qwen3-8b", "object": "model", "owned_by": "organization_owner" });
335 assert!(parse_model_limits(&item).is_none());
336 }
337
338 #[test]
343 fn deepseek_shape_is_parsed() {
344 let item = json!({
345 "id": "deepseek-flash",
346 "object": "model",
347 "owned_by": "deepseek",
348 "context_window": 1_048_576,
349 "max_output_tokens": 393_216
350 });
351 let l = parse_model_limits(&item).unwrap();
352 assert_eq!(l.context_window, Some(1_048_576));
353 assert_eq!(l.max_output, Some(393_216));
354 assert_eq!(l.source, LimitSource::Endpoint);
355 }
356
357 #[test]
360 fn or_tracks_source_per_field() {
361 let user = TokenLimits::from_user(Some(64_000), None);
362 let endpoint = TokenLimits::from_endpoint(None, None);
363 let preset = TokenLimits::from_preset(Some(128_000), Some(8192));
364 let l = user.or(endpoint).or(preset);
365 assert_eq!(l.context_window_source, Some(LimitSource::User));
366 assert_eq!(l.max_output_source, Some(LimitSource::Preset));
367 assert_eq!(
368 l.source,
369 LimitSource::User,
370 "整条的 source 语义不变,兼容已有调用方"
371 );
372
373 let only_window = TokenLimits::from_endpoint(Some(200_000), None);
375 assert_eq!(
376 only_window.context_window_source,
377 Some(LimitSource::Endpoint)
378 );
379 assert_eq!(only_window.max_output_source, None);
380
381 let stored = TokenLimits::with_source(Some(1), Some(2), LimitSource::User);
383 assert_eq!(stored.context_window_source, Some(LimitSource::User));
384 assert_eq!(stored.max_output_source, Some(LimitSource::User));
385
386 let j = serde_json::to_string(&l).unwrap();
388 assert!(j.contains(r#""contextWindowSource":"user""#), "{j}");
389 assert!(j.contains(r#""maxOutputSource":"preset""#), "{j}");
390 }
391
392 #[test]
394 fn or_merges_field_by_field() {
395 let endpoint = TokenLimits::from_endpoint(None, Some(32_000));
396 let preset = TokenLimits::from_preset(Some(128_000), Some(8192));
397 let l = endpoint.or(preset);
398 assert_eq!(l.context_window, Some(128_000), "端点没报窗口,由预置补");
399 assert_eq!(l.max_output, Some(32_000), "端点报了的不被覆盖");
400 assert_eq!(l.source, LimitSource::Endpoint);
401
402 let empty = TokenLimits::from_user(None, None);
403 assert_eq!(
404 empty.or(preset),
405 preset,
406 "空的用户设置不能把来源冒充成 User"
407 );
408 }
409
410 #[test]
412 fn limit_source_roundtrip() {
413 for s in [
414 LimitSource::User,
415 LimitSource::Endpoint,
416 LimitSource::Preset,
417 ] {
418 assert_eq!(LimitSource::parse(s.as_str()), Some(s));
419 assert_eq!(serde_json::to_value(s).unwrap(), s.as_str());
420 }
421 assert_eq!(LimitSource::parse("guess"), None);
422 }
423
424 #[test]
428 fn zero_means_unknown() {
429 let item = json!({ "id": "m", "context_length": 0, "max_output_tokens": 0 });
430 assert!(parse_model_limits(&item).is_none());
431 }
432
433 #[test]
435 fn accepts_stringified_numbers() {
436 let item = json!({ "id": "m", "context_length": "32768" });
437 assert_eq!(
438 parse_model_limits(&item).unwrap().context_window,
439 Some(32_768)
440 );
441 }
442
443 #[test]
445 fn source_is_distinguishable() {
446 assert_eq!(
447 parse_model_limits(&json!({ "context_length": 1 }))
448 .unwrap()
449 .source,
450 LimitSource::Endpoint
451 );
452 assert_eq!(
453 TokenLimits::from_preset(Some(1), None).source,
454 LimitSource::Preset
455 );
456 }
457
458 #[test]
460 fn serializes_camel_case_with_snake_source() {
461 let l = TokenLimits::from_endpoint(Some(128_000), Some(8192));
462 let j = serde_json::to_string(&l).unwrap();
463 assert!(j.contains(r#""contextWindow":128000"#), "{j}");
464 assert!(j.contains(r#""maxOutput":8192"#), "{j}");
465 assert!(j.contains(r#""source":"endpoint""#), "{j}");
467 }
468
469 #[test]
471 fn budget_saturates_instead_of_underflowing() {
472 let l = TokenLimits::from_preset(Some(4096), None);
473 assert_eq!(l.input_budget(8192), Some(0));
474 }
475}