use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum LimitSource {
User,
Endpoint,
Preset,
}
impl LimitSource {
pub const fn as_str(self) -> &'static str {
match self {
LimitSource::User => "user",
LimitSource::Endpoint => "endpoint",
LimitSource::Preset => "preset",
}
}
pub fn parse(s: &str) -> Option<Self> {
match s.trim() {
"user" => Some(LimitSource::User),
"endpoint" => Some(LimitSource::Endpoint),
"preset" => Some(LimitSource::Preset),
_ => None,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
#[serde(rename_all = "camelCase")]
#[non_exhaustive]
pub struct TokenLimits {
pub context_window: Option<u32>,
pub max_output: Option<u32>,
pub source: LimitSource,
pub context_window_source: Option<LimitSource>,
pub max_output_source: Option<LimitSource>,
}
const fn source_of(value: Option<u32>, source: LimitSource) -> Option<LimitSource> {
match value {
Some(_) => Some(source),
None => None,
}
}
impl TokenLimits {
pub const fn from_endpoint(context_window: Option<u32>, max_output: Option<u32>) -> Self {
Self::with_source(context_window, max_output, LimitSource::Endpoint)
}
pub const fn from_preset(context_window: Option<u32>, max_output: Option<u32>) -> Self {
Self::with_source(context_window, max_output, LimitSource::Preset)
}
pub const fn from_user(context_window: Option<u32>, max_output: Option<u32>) -> Self {
Self::with_source(context_window, max_output, LimitSource::User)
}
pub const fn with_source(
context_window: Option<u32>,
max_output: Option<u32>,
source: LimitSource,
) -> Self {
Self {
context_window,
max_output,
source,
context_window_source: source_of(context_window, source),
max_output_source: source_of(max_output, source),
}
}
pub const fn or(self, fallback: TokenLimits) -> TokenLimits {
if self.is_empty() {
return fallback;
}
let (context_window, context_window_source) = match self.context_window {
Some(v) => (Some(v), self.context_window_source),
None => (fallback.context_window, fallback.context_window_source),
};
let (max_output, max_output_source) = match self.max_output {
Some(v) => (Some(v), self.max_output_source),
None => (fallback.max_output, fallback.max_output_source),
};
TokenLimits {
context_window,
max_output,
source: self.source,
context_window_source,
max_output_source,
}
}
pub const fn is_empty(&self) -> bool {
self.context_window.is_none() && self.max_output.is_none()
}
pub fn input_budget(&self, reserve_output: u32) -> Option<u32> {
self.context_window
.map(|w| w.saturating_sub(reserve_output))
}
}
pub const CONTEXT_WINDOW_FIELDS: &[&str] = &[
"context_length",
"context_window",
"max_context_length",
"max_input_tokens",
];
pub const MAX_OUTPUT_FIELDS: &[&str] =
&["max_completion_tokens", "max_output_tokens", "max_tokens"];
pub fn parse_model_limits(item: &serde_json::Value) -> Option<TokenLimits> {
fn num(v: Option<&serde_json::Value>) -> Option<u32> {
let v = v?;
let n = v
.as_u64()
.or_else(|| v.as_f64().filter(|f| *f >= 0.0).map(|f| f as u64))
.or_else(|| v.as_str()?.trim().parse::<u64>().ok())?;
if n == 0 {
return None;
}
u32::try_from(n).ok()
}
let top = item.get("top_provider");
let pick = |keys: &[&str]| -> Option<u32> {
for k in keys {
if let Some(n) = num(top.and_then(|t| t.get(*k))) {
return Some(n);
}
}
keys.iter().find_map(|k| num(item.get(*k)))
};
let context_window = pick(CONTEXT_WINDOW_FIELDS);
let max_output = pick(MAX_OUTPUT_FIELDS);
if context_window.is_none() && max_output.is_none() {
return None;
}
Some(TokenLimits::from_endpoint(context_window, max_output))
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn nested_top_provider_wins() {
let item = json!({
"id": "some/model",
"context_length": 1_048_576,
"top_provider": { "context_length": 131_072, "max_completion_tokens": 32_768 }
});
let l = parse_model_limits(&item).unwrap();
assert_eq!(l.context_window, Some(131_072), "该取服务商实际提供的");
assert_eq!(l.max_output, Some(32_768));
assert_eq!(l.source, LimitSource::Endpoint);
}
#[test]
fn falls_back_to_top_level() {
let item = json!({ "id": "m", "context_length": 200_000 });
let l = parse_model_limits(&item).unwrap();
assert_eq!(l.context_window, Some(200_000));
assert_eq!(l.max_output, None, "没报就是没报,不猜");
}
#[test]
fn bare_openai_shape_yields_none() {
let item = json!({ "id": "qwen3-8b", "object": "model", "owned_by": "organization_owner" });
assert!(parse_model_limits(&item).is_none());
}
#[test]
fn deepseek_shape_is_parsed() {
let item = json!({
"id": "deepseek-flash",
"object": "model",
"owned_by": "deepseek",
"context_window": 1_048_576,
"max_output_tokens": 393_216
});
let l = parse_model_limits(&item).unwrap();
assert_eq!(l.context_window, Some(1_048_576));
assert_eq!(l.max_output, Some(393_216));
assert_eq!(l.source, LimitSource::Endpoint);
}
#[test]
fn or_tracks_source_per_field() {
let user = TokenLimits::from_user(Some(64_000), None);
let endpoint = TokenLimits::from_endpoint(None, None);
let preset = TokenLimits::from_preset(Some(128_000), Some(8192));
let l = user.or(endpoint).or(preset);
assert_eq!(l.context_window_source, Some(LimitSource::User));
assert_eq!(l.max_output_source, Some(LimitSource::Preset));
assert_eq!(
l.source,
LimitSource::User,
"整条的 source 语义不变,兼容已有调用方"
);
let only_window = TokenLimits::from_endpoint(Some(200_000), None);
assert_eq!(
only_window.context_window_source,
Some(LimitSource::Endpoint)
);
assert_eq!(only_window.max_output_source, None);
let stored = TokenLimits::with_source(Some(1), Some(2), LimitSource::User);
assert_eq!(stored.context_window_source, Some(LimitSource::User));
assert_eq!(stored.max_output_source, Some(LimitSource::User));
let j = serde_json::to_string(&l).unwrap();
assert!(j.contains(r#""contextWindowSource":"user""#), "{j}");
assert!(j.contains(r#""maxOutputSource":"preset""#), "{j}");
}
#[test]
fn or_merges_field_by_field() {
let endpoint = TokenLimits::from_endpoint(None, Some(32_000));
let preset = TokenLimits::from_preset(Some(128_000), Some(8192));
let l = endpoint.or(preset);
assert_eq!(l.context_window, Some(128_000), "端点没报窗口,由预置补");
assert_eq!(l.max_output, Some(32_000), "端点报了的不被覆盖");
assert_eq!(l.source, LimitSource::Endpoint);
let empty = TokenLimits::from_user(None, None);
assert_eq!(
empty.or(preset),
preset,
"空的用户设置不能把来源冒充成 User"
);
}
#[test]
fn limit_source_roundtrip() {
for s in [
LimitSource::User,
LimitSource::Endpoint,
LimitSource::Preset,
] {
assert_eq!(LimitSource::parse(s.as_str()), Some(s));
assert_eq!(serde_json::to_value(s).unwrap(), s.as_str());
}
assert_eq!(LimitSource::parse("guess"), None);
}
#[test]
fn zero_means_unknown() {
let item = json!({ "id": "m", "context_length": 0, "max_output_tokens": 0 });
assert!(parse_model_limits(&item).is_none());
}
#[test]
fn accepts_stringified_numbers() {
let item = json!({ "id": "m", "context_length": "32768" });
assert_eq!(
parse_model_limits(&item).unwrap().context_window,
Some(32_768)
);
}
#[test]
fn source_is_distinguishable() {
assert_eq!(
parse_model_limits(&json!({ "context_length": 1 }))
.unwrap()
.source,
LimitSource::Endpoint
);
assert_eq!(
TokenLimits::from_preset(Some(1), None).source,
LimitSource::Preset
);
}
#[test]
fn serializes_camel_case_with_snake_source() {
let l = TokenLimits::from_endpoint(Some(128_000), Some(8192));
let j = serde_json::to_string(&l).unwrap();
assert!(j.contains(r#""contextWindow":128000"#), "{j}");
assert!(j.contains(r#""maxOutput":8192"#), "{j}");
assert!(j.contains(r#""source":"endpoint""#), "{j}");
}
#[test]
fn budget_saturates_instead_of_underflowing() {
let l = TokenLimits::from_preset(Some(4096), None);
assert_eq!(l.input_budget(8192), Some(0));
}
}