Skip to main content

codex_helper_core/provider_signals/
model.rs

1use serde::{Deserialize, Serialize};
2
3use crate::runtime_identity::ProviderEndpointKey;
4
5#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq, PartialOrd, Ord, Default)]
6#[serde(rename_all = "snake_case")]
7pub enum ProviderSignalConfidence {
8    Low,
9    #[default]
10    Medium,
11    High,
12}
13
14#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
15#[serde(rename_all = "snake_case")]
16pub enum ProviderSignalKind {
17    Quota,
18    RateLimit,
19    Capacity,
20    Transport,
21    Balance,
22    ServiceStatus,
23    Capability,
24    LocalConcurrency,
25    #[serde(other)]
26    Unknown,
27}
28
29impl ProviderSignalKind {
30    pub fn code(&self) -> &'static str {
31        match self {
32            ProviderSignalKind::Quota => "quota",
33            ProviderSignalKind::RateLimit => "rate_limit",
34            ProviderSignalKind::Capacity => "capacity",
35            ProviderSignalKind::Transport => "transport",
36            ProviderSignalKind::Balance => "balance",
37            ProviderSignalKind::ServiceStatus => "service_status",
38            ProviderSignalKind::Capability => "capability",
39            ProviderSignalKind::LocalConcurrency => "local_concurrency",
40            ProviderSignalKind::Unknown => "unknown",
41        }
42    }
43}
44
45#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
46#[serde(rename_all = "snake_case")]
47pub enum ProviderSignalSource {
48    UpstreamResponse,
49    ResponseHeaders,
50    BalanceSnapshot,
51    ServiceStatus,
52    CapabilityProbe,
53    LocalScheduler,
54    RouteAttempt,
55}
56
57#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
58#[serde(rename_all = "snake_case")]
59pub enum ProviderSignalTarget {
60    ProviderEndpoint {
61        provider_endpoint_key: ProviderEndpointKey,
62    },
63    Provider {
64        service: String,
65        provider_id: String,
66    },
67    Service {
68        service: String,
69    },
70}
71
72impl ProviderSignalTarget {
73    pub fn provider_endpoint_key(&self) -> Option<&ProviderEndpointKey> {
74        match self {
75            Self::ProviderEndpoint {
76                provider_endpoint_key,
77            } => Some(provider_endpoint_key),
78            Self::Provider { .. } | Self::Service { .. } => None,
79        }
80    }
81}
82
83#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq, Default)]
84pub struct ProviderSignalTrace {
85    #[serde(skip_serializing_if = "Option::is_none")]
86    pub trace_id: Option<String>,
87    #[serde(skip_serializing_if = "Option::is_none")]
88    pub cf_ray: Option<String>,
89    #[serde(skip_serializing_if = "Option::is_none")]
90    pub upstream_request_id: Option<String>,
91}
92
93#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
94pub struct ProviderSignal {
95    pub kind: ProviderSignalKind,
96    #[serde(default, skip_serializing_if = "Option::is_none")]
97    pub code: Option<String>,
98    pub source: ProviderSignalSource,
99    pub target: ProviderSignalTarget,
100    pub confidence: ProviderSignalConfidence,
101    pub observed_at_ms: u64,
102    #[serde(default, skip_serializing_if = "bool_is_false")]
103    pub route_facing: bool,
104    #[serde(skip_serializing_if = "Option::is_none")]
105    pub retry_after_secs: Option<u64>,
106    #[serde(skip_serializing_if = "Option::is_none")]
107    pub reset_after_secs: Option<u64>,
108    #[serde(skip_serializing_if = "Option::is_none")]
109    pub reason: Option<String>,
110    #[serde(skip_serializing_if = "Option::is_none")]
111    pub error_class: Option<String>,
112    #[serde(default, skip_serializing_if = "ProviderSignalTrace::is_empty")]
113    pub trace: ProviderSignalTrace,
114}
115
116impl ProviderSignal {
117    pub fn high_confidence_route_facing(
118        kind: ProviderSignalKind,
119        source: ProviderSignalSource,
120        target: ProviderSignalTarget,
121        observed_at_ms: u64,
122    ) -> Self {
123        Self {
124            code: Some(kind.code().to_string()),
125            kind,
126            source,
127            target,
128            confidence: ProviderSignalConfidence::High,
129            observed_at_ms,
130            route_facing: true,
131            retry_after_secs: None,
132            reset_after_secs: None,
133            reason: None,
134            error_class: None,
135            trace: ProviderSignalTrace::default(),
136        }
137    }
138
139    pub fn cooldown_horizon_secs(&self) -> Option<u64> {
140        self.reset_after_secs.or(self.retry_after_secs)
141    }
142
143    pub fn is_high_confidence_route_facing(&self) -> bool {
144        self.route_facing && self.confidence >= ProviderSignalConfidence::High
145    }
146
147    pub fn stable_code(&self) -> &str {
148        self.code.as_deref().unwrap_or_else(|| self.kind.code())
149    }
150}
151
152impl ProviderSignalTrace {
153    pub fn is_empty(&self) -> bool {
154        self.trace_id.is_none() && self.cf_ray.is_none() && self.upstream_request_id.is_none()
155    }
156}
157
158fn bool_is_false(value: &bool) -> bool {
159    !*value
160}
161
162#[cfg(test)]
163mod tests {
164    use super::*;
165
166    #[test]
167    fn provider_endpoint_target_exposes_stable_key() {
168        let key = ProviderEndpointKey::new("codex", "monthly", "default");
169        let target = ProviderSignalTarget::ProviderEndpoint {
170            provider_endpoint_key: key.clone(),
171        };
172
173        assert_eq!(target.provider_endpoint_key(), Some(&key));
174    }
175
176    #[test]
177    fn cooldown_horizon_prefers_reset_then_retry_after() {
178        let mut signal = ProviderSignal::high_confidence_route_facing(
179            ProviderSignalKind::Quota,
180            ProviderSignalSource::UpstreamResponse,
181            ProviderSignalTarget::ProviderEndpoint {
182                provider_endpoint_key: ProviderEndpointKey::new("codex", "monthly", "default"),
183            },
184            100,
185        );
186        signal.retry_after_secs = Some(30);
187        signal.reset_after_secs = Some(60);
188
189        assert_eq!(signal.cooldown_horizon_secs(), Some(60));
190        assert!(signal.is_high_confidence_route_facing());
191        assert_eq!(signal.stable_code(), "quota");
192    }
193
194    #[test]
195    fn provider_signal_serializes_additive_code() {
196        let signal = ProviderSignal::high_confidence_route_facing(
197            ProviderSignalKind::RateLimit,
198            ProviderSignalSource::UpstreamResponse,
199            ProviderSignalTarget::ProviderEndpoint {
200                provider_endpoint_key: ProviderEndpointKey::new("codex", "monthly", "default"),
201            },
202            100,
203        );
204        let value = serde_json::to_value(&signal).expect("serialize signal");
205
206        assert_eq!(value["kind"].as_str(), Some("rate_limit"));
207        assert_eq!(value["code"].as_str(), Some("rate_limit"));
208    }
209
210    #[test]
211    fn unknown_provider_signal_kind_deserializes_as_unknown_code() {
212        let signal: ProviderSignal = serde_json::from_value(serde_json::json!({
213            "kind": "future_signal",
214            "source": "upstream_response",
215            "target": {
216                "service": { "service": "codex" }
217            },
218            "confidence": "medium",
219            "observed_at_ms": 100
220        }))
221        .expect("deserialize unknown signal kind");
222
223        assert_eq!(signal.kind, ProviderSignalKind::Unknown);
224        assert_eq!(signal.stable_code(), "unknown");
225    }
226}