Skip to main content

atman_runtime/
provider.rs

1use std::collections::HashMap;
2use std::sync::Arc;
3
4use tokio::sync::broadcast;
5use tokio_util::sync::CancellationToken;
6
7use crate::error::RuntimeError;
8use crate::event::{NodeEvent, Observable};
9use crate::message::{Message, MessagePart, MessageRole};
10use crate::tool::BoxFut;
11use crate::value::Value;
12
13#[derive(Debug, Clone)]
14pub struct LlmRequest {
15    pub model: String,
16    pub messages: Vec<Message>,
17    pub system: Option<String>,
18    pub input: Value,
19    pub schema: Option<String>,
20    pub cache_prompt: bool,
21    pub tools: Vec<crate::tool::ToolSpec>,
22    pub thinking_enabled: bool,
23    /// Seconds without a streaming chunk before the call is cancelled and
24    /// retried.  Default 120 s.  0 disables stall detection.
25    pub stall_timeout_secs: u64,
26}
27
28#[derive(Debug, Clone, Default, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
29pub struct TokenUsage {
30    pub input: u64,
31    pub cached_input: u64,
32    pub output: u64,
33    pub cache_write: u64,
34    pub reasoning_tokens: u64,
35}
36
37impl TokenUsage {
38    pub fn total(&self) -> u64 {
39        self.input
40            .saturating_add(self.cached_input)
41            .saturating_add(self.output)
42            .saturating_add(self.cache_write)
43    }
44}
45
46#[derive(Debug, Clone, Default, PartialEq, Eq)]
47pub struct CallTiming {
48    pub total_ms: u64,
49    pub ttft_ms: Option<u64>,
50}
51
52impl CallTiming {
53    pub fn tokens_per_second(&self, output_tokens: u64) -> Option<f64> {
54        let ttft = self.ttft_ms? as f64;
55        let total = self.total_ms as f64;
56        let gen_ms = total - ttft;
57        if gen_ms <= 0.0 || output_tokens == 0 {
58            return None;
59        }
60        Some(output_tokens as f64 / (gen_ms / 1000.0))
61    }
62}
63
64#[derive(Debug, Clone, PartialEq, Eq)]
65pub enum StopReason {
66    End,
67    ToolUse,
68    Length,
69    Cancelled,
70}
71
72#[derive(Debug, Clone)]
73pub struct AssistantMessage {
74    pub message: Message,
75    pub stop_reason: StopReason,
76    pub token_usage: TokenUsage,
77    #[allow(dead_code)]
78    pub timing: CallTiming,
79    pub model: String,
80    pub response_id: Option<String>,
81}
82
83impl AssistantMessage {
84    pub fn text_only(msg: Message) -> Self {
85        Self {
86            message: msg,
87            stop_reason: StopReason::End,
88            token_usage: TokenUsage::default(),
89            timing: CallTiming::default(),
90            model: String::new(),
91            response_id: None,
92        }
93    }
94
95    pub fn text_concat(&self) -> String {
96        self.message.text_concat()
97    }
98}
99
100pub trait Provider: Send + Sync {
101    fn name(&self) -> &str;
102    fn call<'a>(&'a self, req: LlmRequest) -> BoxFut<'a, Result<AssistantMessage, RuntimeError>>;
103    fn call_streaming(&self, req: LlmRequest) -> Observable<AssistantMessage>;
104
105    /// Discover available models from this provider. Default: empty.
106    fn discover_models(&self) -> BoxFut<'static, Vec<DiscoveredModel>> {
107        Box::pin(async { vec![] })
108    }
109}
110
111#[derive(Debug, Clone)]
112pub struct DiscoveredModel {
113    pub slug: String,
114    pub context_budget: Option<u64>,
115    pub thinking: bool,
116}
117
118pub const DEFAULT_STREAM_BUFFER: usize = 1024;
119
120pub fn wrap_call_as_streaming(
121    call_future: BoxFut<'static, Result<AssistantMessage, RuntimeError>>,
122) -> Observable<AssistantMessage> {
123    let (tx, events) = broadcast::channel(DEFAULT_STREAM_BUFFER);
124    let cancel = CancellationToken::new();
125    let cancel_for_task = cancel.clone();
126    let output: BoxFut<'static, Result<AssistantMessage, RuntimeError>> = Box::pin(async move {
127        tokio::select! {
128            biased;
129            _ = cancel_for_task.cancelled() => {
130                let _ = tx.send(NodeEvent::LlmDone { total_tokens: 0 });
131                Err(RuntimeError::Cancelled("call cancelled".into()))
132            }
133            result = call_future => {
134                match &result {
135                    Ok(am) => {
136                        let text = am.text_concat();
137                        if !text.is_empty() {
138                            let _ = tx.send(NodeEvent::LlmChunk {
139                                text: text.clone(),
140                                cumulative_tokens: estimate_tokens(&text),
141                            });
142                        }
143                        let _ = tx.send(NodeEvent::LlmDone { total_tokens: am.token_usage.output });
144                    }
145                    Err(_) => {
146                        let _ = tx.send(NodeEvent::LlmDone { total_tokens: 0 });
147                    }
148                }
149                result
150            }
151        }
152    });
153    Observable {
154        output,
155        events,
156        cancel,
157    }
158}
159
160pub fn estimate_tokens(text: &str) -> u64 {
161    ((text.len() as f64) / 3.5).ceil() as u64
162}
163
164pub fn assistant_message_to_value(am: &AssistantMessage) -> Value {
165    let has_structural_part = am
166        .message
167        .parts
168        .iter()
169        .any(|p| !matches!(p, MessagePart::Text { .. }));
170    if has_structural_part {
171        return Value::Message(am.message.clone());
172    }
173    let text = am.text_concat();
174    if text.is_empty() {
175        return Value::Message(am.message.clone());
176    }
177    match serde_json::from_str::<serde_json::Value>(&text) {
178        Ok(json) => Value::from_json(json),
179        Err(_) => Value::Str(text),
180    }
181}
182
183pub fn user_text_message(text: impl Into<String>) -> Message {
184    Message {
185        role: MessageRole::User,
186        parts: vec![MessagePart::Text { text: text.into() }],
187        turn_id: crate::event::TurnId::now(),
188    }
189}
190
191#[derive(Default, Clone)]
192pub struct ProviderRegistry {
193    providers: HashMap<String, Arc<dyn Provider>>,
194    default: Option<String>,
195}
196
197impl ProviderRegistry {
198    pub fn new() -> Self {
199        Self::default()
200    }
201
202    pub fn register(&mut self, provider: Arc<dyn Provider>) {
203        let name = provider.name().to_string();
204        if self.default.is_none() {
205            self.default = Some(name.clone());
206        }
207        self.providers.insert(name, provider);
208    }
209
210    pub fn set_default(&mut self, name: &str) {
211        if self.providers.contains_key(name) {
212            self.default = Some(name.to_string());
213        }
214    }
215
216    pub fn resolve(&self, model: &str) -> Option<Arc<dyn Provider>> {
217        if let Some(p) = self.providers.get(model) {
218            return Some(p.clone());
219        }
220        if let Some((prefix, _)) = model.split_once('/')
221            && let Some(p) = self.providers.get(prefix)
222        {
223            return Some(p.clone());
224        }
225        if let Some(entry) = crate::model_registry::model_entry(model)
226            && let Some(ref provider_name) = entry.provider
227        {
228            if let Some(p) = self.providers.get(provider_name) {
229                return Some(p.clone());
230            }
231        }
232        if let Some(entry) = crate::model_registry::model_entry(model) {
233            let provider_name = format!("config:{}", entry.model);
234            if let Some(p) = self.providers.get(&provider_name) {
235                return Some(p.clone());
236            }
237            let provider_name = format!("config:{model}");
238            if let Some(p) = self.providers.get(&provider_name) {
239                return Some(p.clone());
240            }
241        }
242        self.default
243            .as_ref()
244            .and_then(|n| self.providers.get(n).cloned())
245    }
246
247    pub fn get(&self, name: &str) -> Option<Arc<dyn Provider>> {
248        self.providers.get(name).cloned()
249    }
250}
251
252#[cfg(test)]
253mod tests {
254    use super::*;
255    use crate::providers::mock::MockProvider;
256
257    /// Helper: build a registry with a "codex" provider and an "openai" default.
258    fn fixture_registry() -> ProviderRegistry {
259        let mut reg = ProviderRegistry::new();
260        let codex = Arc::new(MockProvider::new("codex"));
261        reg.register(codex);
262        let openai = Arc::new(MockProvider::new("openai"));
263        reg.register(openai);
264        reg
265    }
266
267    #[test]
268    fn resolve_prefix_match_codex_slash_model() {
269        // "codex/gpt-5.6-terra" → split '/' → prefix "codex" → found
270        let reg = fixture_registry();
271        let p = reg.resolve("codex/gpt-5.6-terra").expect("should resolve");
272        assert_eq!(p.name(), "codex");
273    }
274
275    #[test]
276    fn resolve_falls_back_to_default_for_unknown() {
277        let reg = fixture_registry();
278        let p = reg
279            .resolve("some-unknown-model")
280            .expect("should fall back to default");
281        // "codex" was registered first, so it's the default.
282        assert_eq!(p.name(), "codex");
283    }
284
285    #[test]
286    fn resolve_model_registry_provider_field_takes_priority() {
287        // Simulate the Codex bootstrap: register model entry with provider="codex",
288        // resolve by model name that has no '/' separator.
289        crate::model_registry::register_model_entries(vec![(
290            "codex-auto-review".into(),
291            crate::model_registry::ModelEntry {
292                model: "codex-auto-review".into(),
293                provider: Some("codex".into()),
294                ..Default::default()
295            },
296        )]);
297
298        let reg = fixture_registry();
299        let p = reg
300            .resolve("codex-auto-review")
301            .expect("should resolve via model registry provider field");
302        assert_eq!(p.name(), "codex");
303    }
304}