Skip to main content

codex_helper_core/policy_actions/
model.rs

1use serde::{Deserialize, Serialize};
2
3use crate::provider_signals::{
4    ProviderSignal, ProviderSignalConfidence, ProviderSignalKind, ProviderSignalTarget,
5};
6use crate::runtime_identity::ProviderEndpointKey;
7
8#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
9#[serde(rename_all = "snake_case")]
10pub enum PolicyActionKind {
11    Cooldown,
12    #[serde(other)]
13    Unknown,
14}
15
16impl PolicyActionKind {
17    pub fn code(&self) -> &'static str {
18        match self {
19            PolicyActionKind::Cooldown => "cooldown",
20            PolicyActionKind::Unknown => "unknown",
21        }
22    }
23}
24
25#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
26#[serde(rename_all = "snake_case")]
27pub enum PolicyActionOwner {
28    CodexHelper,
29}
30
31#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq, Default)]
32#[serde(rename_all = "snake_case")]
33pub enum PolicyActionRecoveryState {
34    #[default]
35    Active,
36}
37
38#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
39pub struct PolicyAction {
40    pub id: String,
41    pub kind: PolicyActionKind,
42    #[serde(default, skip_serializing_if = "Option::is_none")]
43    pub code: Option<String>,
44    pub owner: PolicyActionOwner,
45    pub provider_endpoint_key: ProviderEndpointKey,
46    pub source_signal: ProviderSignal,
47    pub reason: String,
48    pub confidence: ProviderSignalConfidence,
49    pub created_at_ms: u64,
50    pub expires_at_ms: u64,
51    #[serde(default)]
52    pub recovery_state: PolicyActionRecoveryState,
53    #[serde(default)]
54    pub generation: u64,
55}
56
57impl PolicyAction {
58    pub fn cooldown_from_signal(
59        signal: ProviderSignal,
60        created_at_ms: u64,
61        default_cooldown_secs: u64,
62        generation: u64,
63    ) -> Option<Self> {
64        if !signal.is_high_confidence_route_facing() {
65            return None;
66        }
67        if !matches!(
68            signal.kind,
69            ProviderSignalKind::Quota
70                | ProviderSignalKind::RateLimit
71                | ProviderSignalKind::Capacity
72                | ProviderSignalKind::Transport
73                | ProviderSignalKind::Balance
74        ) {
75            return None;
76        }
77        let provider_endpoint_key = match &signal.target {
78            ProviderSignalTarget::ProviderEndpoint {
79                provider_endpoint_key,
80            } => provider_endpoint_key.clone(),
81            ProviderSignalTarget::Provider { .. } | ProviderSignalTarget::Service { .. } => {
82                return None;
83            }
84        };
85        let cooldown_secs = signal.cooldown_horizon_secs().or_else(|| {
86            (matches!(
87                signal.kind,
88                ProviderSignalKind::Capacity | ProviderSignalKind::Transport
89            ) && default_cooldown_secs > 0)
90                .then_some(default_cooldown_secs)
91        })?;
92        if cooldown_secs == 0 {
93            return None;
94        }
95
96        let reason = signal
97            .reason
98            .clone()
99            .or_else(|| signal.error_class.clone())
100            .unwrap_or_else(|| format!("{:?}", signal.kind).to_ascii_lowercase());
101        Some(Self {
102            id: format!(
103                "codex-helper:{}:{}",
104                provider_endpoint_key.stable_key(),
105                created_at_ms
106            ),
107            kind: PolicyActionKind::Cooldown,
108            code: Some(PolicyActionKind::Cooldown.code().to_string()),
109            owner: PolicyActionOwner::CodexHelper,
110            provider_endpoint_key,
111            source_signal: signal.clone(),
112            reason,
113            confidence: signal.confidence,
114            created_at_ms,
115            expires_at_ms: created_at_ms.saturating_add(cooldown_secs.saturating_mul(1000)),
116            recovery_state: PolicyActionRecoveryState::Active,
117            generation,
118        })
119    }
120
121    pub fn is_active_at(&self, now_ms: u64) -> bool {
122        self.recovery_state == PolicyActionRecoveryState::Active && now_ms < self.expires_at_ms
123    }
124
125    pub fn remaining_secs_at(&self, now_ms: u64) -> Option<u64> {
126        self.is_active_at(now_ms)
127            .then(|| self.expires_at_ms.saturating_sub(now_ms).div_ceil(1000))
128    }
129
130    pub fn stable_code(&self) -> &str {
131        self.code.as_deref().unwrap_or_else(|| self.kind.code())
132    }
133}
134
135#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
136pub struct PolicyActionProjection {
137    pub provider_endpoint_key: ProviderEndpointKey,
138    pub active_cooldown: bool,
139    #[serde(default, skip_serializing_if = "Option::is_none")]
140    pub code: Option<String>,
141    #[serde(skip_serializing_if = "Option::is_none")]
142    pub cooldown_remaining_secs: Option<u64>,
143    #[serde(skip_serializing_if = "Option::is_none")]
144    pub reason: Option<String>,
145    #[serde(skip_serializing_if = "Option::is_none")]
146    pub action_id: Option<String>,
147}
148
149impl PolicyActionProjection {
150    pub fn from_action(action: &PolicyAction, now_ms: u64) -> Option<Self> {
151        let cooldown_remaining_secs = action.remaining_secs_at(now_ms)?;
152        Some(Self {
153            provider_endpoint_key: action.provider_endpoint_key.clone(),
154            active_cooldown: matches!(action.kind, PolicyActionKind::Cooldown),
155            code: action
156                .code
157                .clone()
158                .or_else(|| Some(action.kind.code().to_string())),
159            cooldown_remaining_secs: Some(cooldown_remaining_secs),
160            reason: Some(action.reason.clone()),
161            action_id: Some(action.id.clone()),
162        })
163    }
164}
165
166#[cfg(test)]
167mod tests {
168    use super::*;
169    use crate::provider_signals::{ProviderSignalSource, ProviderSignalTarget};
170
171    fn quota_signal(reset_after_secs: Option<u64>) -> ProviderSignal {
172        let mut signal = ProviderSignal::high_confidence_route_facing(
173            ProviderSignalKind::Quota,
174            ProviderSignalSource::UpstreamResponse,
175            ProviderSignalTarget::ProviderEndpoint {
176                provider_endpoint_key: ProviderEndpointKey::new("codex", "monthly", "default"),
177            },
178            100,
179        );
180        signal.reset_after_secs = reset_after_secs;
181        signal.reason = Some("usage_limit_reached".to_string());
182        signal
183    }
184
185    #[test]
186    fn high_confidence_quota_signal_creates_owned_cooldown() {
187        let action = PolicyAction::cooldown_from_signal(quota_signal(Some(30)), 1_000, 0, 7)
188            .expect("cooldown action");
189
190        assert_eq!(action.owner, PolicyActionOwner::CodexHelper);
191        assert_eq!(action.code.as_deref(), Some("cooldown"));
192        assert_eq!(action.expires_at_ms, 31_000);
193        assert_eq!(action.generation, 7);
194        assert!(action.is_active_at(30_999));
195        assert!(!action.is_active_at(31_000));
196    }
197
198    #[test]
199    fn quota_without_horizon_is_recorded_only() {
200        assert!(PolicyAction::cooldown_from_signal(quota_signal(None), 1_000, 0, 1).is_none());
201    }
202
203    #[test]
204    fn policy_action_projection_serializes_code() {
205        let action = PolicyAction::cooldown_from_signal(quota_signal(Some(30)), 1_000, 0, 7)
206            .expect("cooldown action");
207        let projection = PolicyActionProjection::from_action(&action, 2_000).expect("projection");
208        let value = serde_json::to_value(&projection).expect("serialize projection");
209
210        assert_eq!(value["code"].as_str(), Some("cooldown"));
211        assert_eq!(value["active_cooldown"].as_bool(), Some(true));
212    }
213
214    #[test]
215    fn unknown_policy_action_kind_deserializes_as_unknown_code() {
216        let kind: PolicyActionKind =
217            serde_json::from_str("\"future_action\"").expect("deserialize action kind");
218
219        assert_eq!(kind, PolicyActionKind::Unknown);
220        assert_eq!(kind.code(), "unknown");
221    }
222}