Skip to main content

atman_runtime/
provider.rs

1use std::collections::HashMap;
2use std::fmt;
3use std::str::FromStr;
4use std::sync::Arc;
5
6use serde::{Deserialize, Deserializer, Serialize, Serializer};
7use tokio::sync::broadcast;
8use tokio_util::sync::CancellationToken;
9
10use crate::error::RuntimeError;
11use crate::event::{NodeEvent, Observable};
12use crate::message::{Message, MessageOrigin, MessagePart, MessageRole};
13use crate::tool::BoxFut;
14use crate::value::Value;
15
16#[derive(Debug, Clone, PartialEq, Eq, Hash)]
17pub enum ReasoningEffort {
18    None,
19    Minimal,
20    Low,
21    Medium,
22    High,
23    XHigh,
24    Max,
25    Ultra,
26    Persistent,
27    Custom(String),
28}
29
30impl ReasoningEffort {
31    pub fn as_str(&self) -> &str {
32        match self {
33            Self::None => "none",
34            Self::Minimal => "minimal",
35            Self::Low => "low",
36            Self::Medium => "medium",
37            Self::High => "high",
38            Self::XHigh => "xhigh",
39            Self::Max => "max",
40            Self::Ultra => "ultra",
41            Self::Persistent => "persistent",
42            Self::Custom(value) => value,
43        }
44    }
45}
46
47impl fmt::Display for ReasoningEffort {
48    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
49        f.write_str(self.as_str())
50    }
51}
52
53impl FromStr for ReasoningEffort {
54    type Err = String;
55
56    fn from_str(value: &str) -> Result<Self, Self::Err> {
57        match value.trim().to_ascii_lowercase().as_str() {
58            "none" | "off" | "disabled" => Ok(Self::None),
59            "minimal" => Ok(Self::Minimal),
60            "low" => Ok(Self::Low),
61            "medium" => Ok(Self::Medium),
62            "high" => Ok(Self::High),
63            "xhigh" => Ok(Self::XHigh),
64            "max" => Ok(Self::Max),
65            "ultra" => Ok(Self::Ultra),
66            "persistent" => Ok(Self::Persistent),
67            "" => Err("reasoning effort must not be empty".into()),
68            other => Ok(Self::Custom(other.to_string())),
69        }
70    }
71}
72
73impl Serialize for ReasoningEffort {
74    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
75    where
76        S: Serializer,
77    {
78        serializer.serialize_str(self.as_str())
79    }
80}
81
82impl<'de> Deserialize<'de> for ReasoningEffort {
83    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
84    where
85        D: Deserializer<'de>,
86    {
87        String::deserialize(deserializer)?
88            .parse()
89            .map_err(serde::de::Error::custom)
90    }
91}
92
93#[derive(Debug, Clone, PartialEq, Eq, Hash)]
94pub enum ReasoningExecutionMode {
95    Standard,
96    Pro,
97    Custom(String),
98}
99
100impl ReasoningExecutionMode {
101    pub fn as_str(&self) -> &str {
102        match self {
103            Self::Standard => "standard",
104            Self::Pro => "pro",
105            Self::Custom(value) => value,
106        }
107    }
108}
109
110impl fmt::Display for ReasoningExecutionMode {
111    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
112        f.write_str(self.as_str())
113    }
114}
115
116impl FromStr for ReasoningExecutionMode {
117    type Err = String;
118
119    fn from_str(value: &str) -> Result<Self, Self::Err> {
120        match value.trim().to_ascii_lowercase().as_str() {
121            "standard" => Ok(Self::Standard),
122            "pro" => Ok(Self::Pro),
123            "" => Err("reasoning mode must not be empty".into()),
124            other => Ok(Self::Custom(other.to_string())),
125        }
126    }
127}
128
129impl Serialize for ReasoningExecutionMode {
130    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
131    where
132        S: Serializer,
133    {
134        serializer.serialize_str(self.as_str())
135    }
136}
137
138impl<'de> Deserialize<'de> for ReasoningExecutionMode {
139    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
140    where
141        D: Deserializer<'de>,
142    {
143        String::deserialize(deserializer)?
144            .parse()
145            .map_err(serde::de::Error::custom)
146    }
147}
148
149#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)]
150#[serde(tag = "type", rename_all = "snake_case")]
151pub enum ReasoningSelection {
152    #[default]
153    ProviderDefault,
154    Disabled,
155    Auto {
156        #[serde(default, skip_serializing_if = "Option::is_none")]
157        execution_mode: Option<ReasoningExecutionMode>,
158    },
159    Effort {
160        effort: ReasoningEffort,
161        #[serde(default, skip_serializing_if = "Option::is_none")]
162        execution_mode: Option<ReasoningExecutionMode>,
163    },
164    BudgetTokens {
165        tokens: u32,
166    },
167}
168
169#[derive(Debug, Clone, Copy, PartialEq, Eq)]
170pub enum ReasoningWireProfile {
171    OpenAiOfficial,
172    CompatibleThinking,
173    CodexResponses,
174    AnthropicMessages,
175    Unknown,
176}
177
178const OPENAI_REASONING_EFFORTS: &[ReasoningEffort] = &[
179    ReasoningEffort::Minimal,
180    ReasoningEffort::Low,
181    ReasoningEffort::Medium,
182    ReasoningEffort::High,
183    ReasoningEffort::XHigh,
184    ReasoningEffort::Max,
185    ReasoningEffort::Ultra,
186];
187
188const CODEX_REASONING_EFFORTS: &[ReasoningEffort] = &[
189    ReasoningEffort::Minimal,
190    ReasoningEffort::Low,
191    ReasoningEffort::Medium,
192    ReasoningEffort::High,
193    ReasoningEffort::XHigh,
194    ReasoningEffort::Max,
195    ReasoningEffort::Ultra,
196    ReasoningEffort::Persistent,
197];
198
199const ANTHROPIC_REASONING_EFFORTS: &[ReasoningEffort] = &[
200    ReasoningEffort::Low,
201    ReasoningEffort::Medium,
202    ReasoningEffort::High,
203    ReasoningEffort::Max,
204];
205
206impl ReasoningWireProfile {
207    pub fn fallback_efforts(self) -> &'static [ReasoningEffort] {
208        match self {
209            Self::OpenAiOfficial => OPENAI_REASONING_EFFORTS,
210            Self::CodexResponses => CODEX_REASONING_EFFORTS,
211            Self::AnthropicMessages => ANTHROPIC_REASONING_EFFORTS,
212            Self::CompatibleThinking | Self::Unknown => &[],
213        }
214    }
215
216    pub fn supports_token_budget(self) -> bool {
217        matches!(self, Self::AnthropicMessages)
218    }
219
220    pub fn validate(
221        self,
222        selection: &ReasoningSelection,
223        max_tokens: Option<u32>,
224    ) -> Result<(), String> {
225        if selection.execution_mode().is_some()
226            && !matches!(self, Self::CodexResponses | Self::Unknown)
227        {
228            return Err(match self {
229                Self::AnthropicMessages => {
230                    "Anthropic does not support reasoning execution mode".into()
231                }
232                Self::OpenAiOfficial | Self::CompatibleThinking => {
233                    "Chat Completions does not support reasoning execution mode".into()
234                }
235                Self::CodexResponses | Self::Unknown => unreachable!(),
236            });
237        }
238
239        match (self, selection) {
240            (Self::Unknown, _) => Ok(()),
241            (
242                Self::OpenAiOfficial | Self::CompatibleThinking | Self::CodexResponses,
243                ReasoningSelection::BudgetTokens { .. },
244            ) => Err(match self {
245                Self::CodexResponses => {
246                    "Codex Responses does not support token-budget reasoning".into()
247                }
248                _ => "this OpenAI adapter does not support token-budget reasoning".into(),
249            }),
250            (
251                Self::CompatibleThinking,
252                ReasoningSelection::Effort {
253                    effort: ReasoningEffort::None,
254                    ..
255                },
256            ) => Ok(()),
257            (Self::CompatibleThinking, ReasoningSelection::Effort { effort, .. }) => Err(format!(
258                "compatible thinking profile cannot represent effort `{effort}`; use `auto` or select the official OpenAI profile"
259            )),
260            (
261                Self::AnthropicMessages,
262                ReasoningSelection::Effort {
263                    effort:
264                        ReasoningEffort::Minimal
265                        | ReasoningEffort::XHigh
266                        | ReasoningEffort::Ultra
267                        | ReasoningEffort::Persistent,
268                    ..
269                },
270            ) => Err(format!(
271                "Anthropic Messages cannot represent effort `{}`; use one of: low, medium, high, max",
272                selection.effort().expect("matched effort")
273            )),
274            (Self::AnthropicMessages, ReasoningSelection::BudgetTokens { tokens })
275                if *tokens < 1024 =>
276            {
277                Err("Anthropic thinking budget must be at least 1024 tokens".into())
278            }
279            (Self::AnthropicMessages, ReasoningSelection::BudgetTokens { tokens })
280                if max_tokens.is_some_and(|max| *tokens >= max) =>
281            {
282                Err(format!(
283                    "Anthropic thinking budget ({tokens}) must be lower than max_tokens ({})",
284                    max_tokens.expect("checked max_tokens")
285                ))
286            }
287            _ => Ok(()),
288        }
289    }
290}
291
292impl ReasoningSelection {
293    pub fn enabled(&self) -> bool {
294        !matches!(self, Self::ProviderDefault | Self::Disabled)
295    }
296
297    pub fn effort(&self) -> Option<&ReasoningEffort> {
298        match self {
299            Self::Effort { effort, .. } => Some(effort),
300            _ => None,
301        }
302    }
303
304    pub fn execution_mode(&self) -> Option<&ReasoningExecutionMode> {
305        match self {
306            Self::Auto { execution_mode } | Self::Effort { execution_mode, .. } => {
307                execution_mode.as_ref()
308            }
309            _ => None,
310        }
311    }
312}
313
314impl std::fmt::Display for ReasoningSelection {
315    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
316        match self {
317            Self::ProviderDefault => f.write_str("default"),
318            Self::Disabled => f.write_str("off"),
319            Self::Auto { execution_mode } => {
320                f.write_str("auto")?;
321                if let Some(mode) = execution_mode {
322                    write!(f, "@{mode}")?;
323                }
324                Ok(())
325            }
326            Self::Effort {
327                effort,
328                execution_mode,
329            } => {
330                effort.fmt(f)?;
331                if let Some(mode) = execution_mode {
332                    write!(f, "@{mode}")?;
333                }
334                Ok(())
335            }
336            Self::BudgetTokens { tokens } => write!(f, "budget:{tokens}"),
337        }
338    }
339}
340
341impl std::str::FromStr for ReasoningSelection {
342    type Err = String;
343
344    fn from_str(value: &str) -> Result<Self, Self::Err> {
345        let value = value.trim().to_ascii_lowercase();
346        if matches!(value.as_str(), "default" | "provider_default") {
347            return Ok(Self::ProviderDefault);
348        }
349        if matches!(value.as_str(), "off" | "disabled" | "none") {
350            return Ok(Self::Disabled);
351        }
352        if let Some(tokens) = value.strip_prefix("budget:") {
353            let tokens: u32 = tokens
354                .parse()
355                .map_err(|_| format!("invalid reasoning token budget `{tokens}`"))?;
356            if tokens == 0 {
357                return Err("reasoning token budget must be positive".into());
358            }
359            return Ok(Self::BudgetTokens { tokens });
360        }
361        let (level, execution_mode) = match value.split_once('@') {
362            Some((level, mode)) => (level, Some(mode.parse()?)),
363            None => (value.as_str(), None),
364        };
365        if level == "auto" {
366            return Ok(Self::Auto { execution_mode });
367        }
368        Ok(Self::Effort {
369            effort: level.parse()?,
370            execution_mode,
371        })
372    }
373}
374
375#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize, Default)]
376#[serde(rename_all = "lowercase")]
377pub enum InputModality {
378    #[default]
379    Text,
380    Image,
381    Audio,
382}
383
384#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize, Default)]
385#[serde(rename_all = "lowercase")]
386pub enum ImageDetail {
387    #[default]
388    Auto,
389    Low,
390    High,
391    Original,
392}
393
394impl ImageDetail {
395    pub fn as_str(self) -> &'static str {
396        match self {
397            Self::Auto => "auto",
398            Self::Low => "low",
399            Self::High => "high",
400            Self::Original => "original",
401        }
402    }
403}
404
405#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)]
406pub struct ModelCapabilities {
407    #[serde(default)]
408    pub reasoning_efforts: Vec<ReasoningEffort>,
409    #[serde(default, skip_serializing_if = "Option::is_none")]
410    pub default_reasoning_effort: Option<ReasoningEffort>,
411    #[serde(default)]
412    pub reasoning_modes: Vec<ReasoningExecutionMode>,
413    #[serde(default, skip_serializing_if = "Option::is_none")]
414    pub default_reasoning_mode: Option<ReasoningExecutionMode>,
415    #[serde(default)]
416    pub input_modalities: Vec<InputModality>,
417}
418
419/// Wire capabilities of a provider endpoint.
420///
421/// Defaults are deliberately conservative because OpenAI-compatible endpoints
422/// do not necessarily implement optional fields from the official API.
423#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
424pub struct ProviderCapabilities {
425    pub prompt_cache_key: bool,
426    pub context_prefix_profile: crate::context_plan::ContextPrefixProfile,
427}
428
429#[derive(Debug, Clone)]
430pub struct LlmRequest {
431    pub model: String,
432    pub messages: Vec<Message>,
433    pub system: Option<String>,
434    pub input: Value,
435    pub schema: Option<String>,
436    pub cache_prompt: bool,
437    pub prompt_cache_key: Option<String>,
438    pub tools: Vec<crate::tool::ToolSpec>,
439    pub reasoning: ReasoningSelection,
440    /// Seconds without a streaming chunk before the call is cancelled and
441    /// retried.  Default 120 s.  0 disables stall detection.
442    pub stall_timeout_secs: u64,
443}
444
445#[derive(Debug, Clone, Default, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
446#[serde(default)]
447pub struct TokenUsage {
448    /// Regular input tokens that were neither read from nor written to cache.
449    pub input: u64,
450    /// Input tokens read from cache.
451    pub cached_input: u64,
452    pub output: u64,
453    /// Input tokens written to cache. This lane is disjoint from `input`.
454    pub cache_write: u64,
455    pub reasoning_tokens: u64,
456}
457
458impl TokenUsage {
459    pub fn prompt_input(&self) -> u64 {
460        self.input
461            .saturating_add(self.cached_input)
462            .saturating_add(self.cache_write)
463    }
464
465    pub fn total(&self) -> u64 {
466        self.prompt_input().saturating_add(self.output)
467    }
468}
469
470pub(crate) fn regular_input_tokens(total_input: u64, cached_input: u64, cache_write: u64) -> u64 {
471    total_input
472        .saturating_sub(cached_input)
473        .saturating_sub(cache_write)
474}
475
476#[derive(Debug, Clone, Default, PartialEq, Eq)]
477pub struct CallTiming {
478    pub total_ms: u64,
479    pub ttft_ms: Option<u64>,
480}
481
482impl CallTiming {
483    pub fn tokens_per_second(&self, output_tokens: u64) -> Option<f64> {
484        let ttft = self.ttft_ms? as f64;
485        let total = self.total_ms as f64;
486        let gen_ms = total - ttft;
487        if gen_ms <= 0.0 || output_tokens == 0 {
488            return None;
489        }
490        Some(output_tokens as f64 / (gen_ms / 1000.0))
491    }
492}
493
494#[derive(Debug, Clone, PartialEq, Eq)]
495pub enum StopReason {
496    End,
497    ToolUse,
498    Length,
499    Cancelled,
500}
501
502#[derive(Debug, Clone)]
503pub struct AssistantMessage {
504    pub message: Message,
505    pub stop_reason: StopReason,
506    pub token_usage: TokenUsage,
507    #[allow(dead_code)]
508    pub timing: CallTiming,
509    pub model: String,
510    pub response_id: Option<String>,
511}
512
513impl AssistantMessage {
514    pub fn text_only(msg: Message) -> Self {
515        Self {
516            message: msg,
517            stop_reason: StopReason::End,
518            token_usage: TokenUsage::default(),
519            timing: CallTiming::default(),
520            model: String::new(),
521            response_id: None,
522        }
523    }
524
525    pub fn text_concat(&self) -> String {
526        self.message.text_concat()
527    }
528}
529
530pub(crate) fn bounded_utf8_prefix(value: &str, max_bytes: usize) -> &str {
531    let mut end = value.len().min(max_bytes);
532    while !value.is_char_boundary(end) {
533        end -= 1;
534    }
535    &value[..end]
536}
537
538pub trait Provider: Send + Sync {
539    fn name(&self) -> &str;
540    fn capabilities(&self) -> ProviderCapabilities {
541        ProviderCapabilities::default()
542    }
543    fn call<'a>(&'a self, req: LlmRequest) -> BoxFut<'a, Result<AssistantMessage, RuntimeError>>;
544    fn call_streaming(&self, req: LlmRequest) -> Observable<AssistantMessage>;
545
546    /// Project the cacheable prompt sequence in provider order.
547    ///
548    /// The default preserves provider-neutral tool/system/message boundaries.
549    /// Providers with a distinct wire projection override this method.
550    fn context_prefix(
551        &self,
552        req: &LlmRequest,
553    ) -> Result<crate::context_plan::ContextPrefixSnapshot, RuntimeError> {
554        crate::context_plan::ContextPrefixSnapshot::provider_neutral(req)
555    }
556
557    /// Discover available models without capability provenance.
558    fn discover_models(&self) -> BoxFut<'static, Vec<DiscoveredModel>> {
559        Box::pin(async { vec![] })
560    }
561
562    /// Discover available models with capability provenance and typed failures.
563    /// The default treats non-empty compatibility results as legacy capability
564    /// data and reports an empty result as unsupported.
565    fn try_discover_models(
566        &self,
567    ) -> BoxFut<'static, Result<Vec<DiscoveredModelDetails>, ModelDiscoveryError>> {
568        let discovery = self.discover_models();
569        Box::pin(async move {
570            let models = discovery.await;
571            if models.is_empty() {
572                return Err(ModelDiscoveryError::Unsupported);
573            }
574            Ok(models
575                .into_iter()
576                .map(DiscoveredModelDetails::from)
577                .collect())
578        })
579    }
580
581    fn test_connection(&self) -> BoxFut<'_, Result<String, String>> {
582        Box::pin(async { Err("test_connection not implemented".into()) })
583    }
584}
585
586#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
587#[non_exhaustive]
588pub enum ModelDiscoveryError {
589    #[error("model discovery is not supported by this provider")]
590    Unsupported,
591    #[error("model discovery transport failed: {0}")]
592    Transport(String),
593    #[error("model discovery returned HTTP {status}: {body}")]
594    Http { status: u16, body: String },
595    #[error("model discovery returned an invalid response: {0}")]
596    InvalidResponse(String),
597}
598
599#[derive(Debug, Clone)]
600pub struct DiscoveredModel {
601    pub slug: String,
602    pub context_budget: Option<u64>,
603    pub thinking: bool,
604}
605
606#[derive(Debug, Clone, PartialEq, Eq)]
607pub struct DiscoveredModelDetails {
608    pub slug: String,
609    pub context_budget: Option<u64>,
610    pub capability_knowledge: CapabilityKnowledge,
611}
612
613#[derive(Debug, Clone, PartialEq, Eq)]
614#[non_exhaustive]
615pub enum CapabilityKnowledge {
616    Legacy { thinking: bool },
617    Advertised(ModelCapabilities),
618}
619
620impl CapabilityKnowledge {
621    pub fn thinking(&self) -> bool {
622        match self {
623            Self::Legacy { thinking } => *thinking,
624            Self::Advertised(capabilities) => {
625                capabilities
626                    .reasoning_efforts
627                    .iter()
628                    .chain(capabilities.default_reasoning_effort.iter())
629                    .any(|effort| !matches!(effort, ReasoningEffort::None))
630                    || !capabilities.reasoning_modes.is_empty()
631                    || capabilities.default_reasoning_mode.is_some()
632            }
633        }
634    }
635
636    pub fn advertised(&self) -> Option<&ModelCapabilities> {
637        match self {
638            Self::Legacy { .. } => None,
639            Self::Advertised(capabilities) => Some(capabilities),
640        }
641    }
642}
643
644impl From<DiscoveredModel> for DiscoveredModelDetails {
645    fn from(model: DiscoveredModel) -> Self {
646        Self {
647            slug: model.slug,
648            context_budget: model.context_budget,
649            capability_knowledge: CapabilityKnowledge::Legacy {
650                thinking: model.thinking,
651            },
652        }
653    }
654}
655
656impl From<DiscoveredModelDetails> for DiscoveredModel {
657    fn from(model: DiscoveredModelDetails) -> Self {
658        Self {
659            slug: model.slug,
660            context_budget: model.context_budget,
661            thinking: model.capability_knowledge.thinking(),
662        }
663    }
664}
665
666pub const DEFAULT_STREAM_BUFFER: usize = 1024;
667
668pub fn wrap_call_as_streaming(
669    call_future: BoxFut<'static, Result<AssistantMessage, RuntimeError>>,
670) -> Observable<AssistantMessage> {
671    let (tx, events) = broadcast::channel(DEFAULT_STREAM_BUFFER);
672    let cancel = CancellationToken::new();
673    let cancel_for_task = cancel.clone();
674    let output: BoxFut<'static, Result<AssistantMessage, RuntimeError>> = Box::pin(async move {
675        tokio::select! {
676            biased;
677            _ = cancel_for_task.cancelled() => {
678                let _ = tx.send(NodeEvent::LlmDone { total_tokens: 0 });
679                Err(RuntimeError::Cancelled("call cancelled".into()))
680            }
681            result = call_future => {
682                match &result {
683                    Ok(am) => {
684                        let text = am.text_concat();
685                        if !text.is_empty() {
686                            let _ = tx.send(NodeEvent::LlmChunk {
687                                text: text.clone(),
688                                cumulative_tokens: estimate_tokens(&text),
689                            });
690                        }
691                        let _ = tx.send(NodeEvent::LlmDone { total_tokens: am.token_usage.output });
692                    }
693                    Err(_) => {
694                        let _ = tx.send(NodeEvent::LlmDone { total_tokens: 0 });
695                    }
696                }
697                result
698            }
699        }
700    });
701    Observable {
702        output,
703        events,
704        cancel,
705    }
706}
707
708pub fn estimate_tokens(text: &str) -> u64 {
709    ((text.len() as f64) / 3.5).ceil() as u64
710}
711
712pub fn assistant_message_to_value(am: &AssistantMessage) -> Value {
713    let has_structural_part = am
714        .message
715        .parts
716        .iter()
717        .any(|p| !matches!(p, MessagePart::Text { .. }));
718    if has_structural_part {
719        return Value::Message(am.message.clone());
720    }
721    let text = am.text_concat();
722    if text.is_empty() {
723        return Value::Message(am.message.clone());
724    }
725    match serde_json::from_str::<serde_json::Value>(&text) {
726        Ok(json) => Value::from_json(json),
727        Err(_) => Value::Str(text),
728    }
729}
730
731pub fn user_text_message(text: impl Into<String>) -> Message {
732    Message {
733        role: MessageRole::User,
734        parts: vec![MessagePart::Text { text: text.into() }],
735        turn_id: crate::event::TurnId::now(),
736        origin: MessageOrigin::User,
737    }
738}
739
740#[derive(Default, Clone)]
741pub struct ProviderRegistry {
742    providers: std::sync::Arc<std::sync::RwLock<HashMap<String, Arc<dyn Provider>>>>,
743    default: std::sync::Arc<std::sync::RwLock<Option<String>>>,
744    lifecycle_owner:
745        std::sync::Arc<std::sync::Mutex<Option<crate::provider_lifecycle::ProviderLifecycleOwner>>>,
746}
747
748#[derive(Clone)]
749pub(crate) struct WeakProviderRegistry {
750    providers: std::sync::Weak<std::sync::RwLock<HashMap<String, Arc<dyn Provider>>>>,
751    default: std::sync::Weak<std::sync::RwLock<Option<String>>>,
752    lifecycle_owner: std::sync::Weak<
753        std::sync::Mutex<Option<crate::provider_lifecycle::ProviderLifecycleOwner>>,
754    >,
755}
756
757impl ProviderRegistry {
758    pub fn new() -> Self {
759        Self::default()
760    }
761
762    pub fn register(&self, provider: Arc<dyn Provider>) {
763        let name = provider.name().to_string();
764        drop(self.register_named(name, provider));
765    }
766
767    pub(crate) fn register_named(
768        &self,
769        name: String,
770        provider: Arc<dyn Provider>,
771    ) -> Option<Arc<dyn Provider>> {
772        let mut providers = self
773            .providers
774            .write()
775            .unwrap_or_else(std::sync::PoisonError::into_inner);
776        let mut defaults = self
777            .default
778            .write()
779            .unwrap_or_else(std::sync::PoisonError::into_inner);
780        if defaults.is_none() {
781            *defaults = Some(name.clone());
782        }
783        providers.insert(name, provider)
784    }
785
786    pub fn set_default(&self, name: &str) {
787        let providers = self.providers.read().unwrap();
788        if providers.contains_key(name) {
789            *self.default.write().unwrap() = Some(name.to_string());
790        }
791    }
792
793    pub fn remove(&self, name: &str) -> bool {
794        self.take_named(name).is_some()
795    }
796
797    pub(crate) fn take_named(&self, name: &str) -> Option<Arc<dyn Provider>> {
798        let mut providers = self.providers.write().unwrap();
799        let removed = providers.remove(name);
800        if removed.is_some() {
801            let mut default = self.default.write().unwrap();
802            if default.as_deref() == Some(name) {
803                *default = None;
804            }
805        }
806        removed
807    }
808
809    pub fn contains(&self, name: &str) -> bool {
810        self.providers.read().unwrap().contains_key(name)
811    }
812
813    pub(crate) fn shares_storage_with(&self, other: &Self) -> bool {
814        Arc::ptr_eq(&self.providers, &other.providers)
815    }
816
817    pub(crate) fn attach_provider_lifecycle(
818        &self,
819        hub: crate::config_hub::ConfigHub,
820    ) -> Option<crate::provider_lifecycle::ProviderLifecycle> {
821        let mut owner = self
822            .lifecycle_owner
823            .lock()
824            .unwrap_or_else(std::sync::PoisonError::into_inner);
825        if owner.is_some() {
826            return None;
827        }
828        let (lifecycle, replaced) =
829            crate::provider_lifecycle::ProviderLifecycle::new_deferred(hub, self.clone());
830        *owner = Some(lifecycle.owner());
831        drop(owner);
832        drop(replaced);
833        Some(lifecycle)
834    }
835
836    pub(crate) fn provider_lifecycle(
837        &self,
838    ) -> Option<crate::provider_lifecycle::ProviderLifecycle> {
839        let owner = self
840            .lifecycle_owner
841            .lock()
842            .unwrap_or_else(std::sync::PoisonError::into_inner)
843            .clone()?;
844        Some(crate::provider_lifecycle::ProviderLifecycle::from_owner(
845            owner,
846            self.clone(),
847        ))
848    }
849
850    pub(crate) fn downgrade(&self) -> WeakProviderRegistry {
851        WeakProviderRegistry {
852            providers: Arc::downgrade(&self.providers),
853            default: Arc::downgrade(&self.default),
854            lifecycle_owner: Arc::downgrade(&self.lifecycle_owner),
855        }
856    }
857
858    pub fn resolve(&self, model: &str) -> Option<Arc<dyn Provider>> {
859        let providers = self.providers.read().unwrap();
860        if let Some(p) = providers.get(model) {
861            return Some(p.clone());
862        }
863        if let Some(entry) = crate::model_registry::model_entry(model)
864            && let Some(ref provider_name) = entry.provider
865        {
866            if let Some(provider) = providers.get(provider_name) {
867                return Some(provider.clone());
868            }
869            if !crate::model_registry::is_provider_enabled(provider_name) {
870                return None;
871            }
872            let config_key = format!("config:{provider_name}");
873            return providers.get(&config_key).cloned();
874        }
875        if let Some((prefix, _)) = model.split_once('/')
876            && let Some(p) = providers.get(prefix)
877        {
878            return Some(p.clone());
879        }
880        None
881    }
882
883    pub fn get(&self, name: &str) -> Option<Arc<dyn Provider>> {
884        self.providers.read().unwrap().get(name).cloned()
885    }
886}
887
888impl WeakProviderRegistry {
889    pub(crate) fn upgrade(&self) -> Option<ProviderRegistry> {
890        Some(ProviderRegistry {
891            providers: self.providers.upgrade()?,
892            default: self.default.upgrade()?,
893            lifecycle_owner: self.lifecycle_owner.upgrade()?,
894        })
895    }
896}
897
898#[cfg(test)]
899mod tests {
900    use super::*;
901    use crate::providers::mock::MockProvider;
902    use std::sync::atomic::{AtomicBool, Ordering};
903
904    struct LegacyDiscoveryProvider;
905
906    #[test]
907    fn bounded_utf8_prefix_never_splits_a_character() {
908        let value = format!("{}z", "界".repeat(67));
909        let prefix = bounded_utf8_prefix(&value, 200);
910        assert!(prefix.len() <= 200);
911        assert_eq!(prefix, "界".repeat(66));
912    }
913
914    #[test]
915    fn token_usage_prompt_lanes_are_disjoint() {
916        let usage = TokenUsage {
917            input: 20,
918            cached_input: 80,
919            output: 10,
920            cache_write: 50,
921            reasoning_tokens: 0,
922        };
923
924        assert_eq!(usage.prompt_input(), 150);
925        assert_eq!(usage.total(), 160);
926        assert_eq!(regular_input_tokens(150, 80, 50), 20);
927    }
928
929    struct OwnerLockProbeProvider {
930        name: String,
931        lifecycle_owner:
932            Arc<std::sync::Mutex<Option<crate::provider_lifecycle::ProviderLifecycleOwner>>>,
933        owner_was_unlocked: Arc<AtomicBool>,
934    }
935
936    impl Drop for OwnerLockProbeProvider {
937        fn drop(&mut self) {
938            self.owner_was_unlocked
939                .store(self.lifecycle_owner.try_lock().is_ok(), Ordering::SeqCst);
940        }
941    }
942
943    impl Provider for OwnerLockProbeProvider {
944        fn name(&self) -> &str {
945            &self.name
946        }
947
948        fn call<'a>(
949            &'a self,
950            _req: LlmRequest,
951        ) -> BoxFut<'a, Result<AssistantMessage, RuntimeError>> {
952            Box::pin(async { unreachable!("not used by owner lock test") })
953        }
954
955        fn call_streaming(&self, _req: LlmRequest) -> Observable<AssistantMessage> {
956            unreachable!("not used by owner lock test")
957        }
958    }
959
960    impl Provider for LegacyDiscoveryProvider {
961        fn name(&self) -> &str {
962            "legacy"
963        }
964
965        fn call<'a>(
966            &'a self,
967            _req: LlmRequest,
968        ) -> BoxFut<'a, Result<AssistantMessage, RuntimeError>> {
969            Box::pin(async { unreachable!("not used by discovery test") })
970        }
971
972        fn call_streaming(&self, _req: LlmRequest) -> Observable<AssistantMessage> {
973            unreachable!("not used by discovery test")
974        }
975
976        fn discover_models(&self) -> BoxFut<'static, Vec<DiscoveredModel>> {
977            Box::pin(async {
978                vec![DiscoveredModel {
979                    slug: "legacy/model".into(),
980                    context_budget: Some(8_192),
981                    thinking: true,
982                }]
983            })
984        }
985    }
986
987    fn fixture_registry() -> ProviderRegistry {
988        let reg = ProviderRegistry::new();
989        let codex = Arc::new(MockProvider::new("codex"));
990        reg.register(codex);
991        let openai = Arc::new(MockProvider::new("openai"));
992        reg.register(openai);
993        reg
994    }
995
996    #[tokio::test]
997    async fn fallible_discovery_adapts_legacy_provider_implementations() {
998        let models = LegacyDiscoveryProvider.try_discover_models().await.unwrap();
999
1000        assert_eq!(models.len(), 1);
1001        assert_eq!(models[0].slug, "legacy/model");
1002        assert_eq!(models[0].context_budget, Some(8_192));
1003        assert_eq!(
1004            models[0].capability_knowledge,
1005            CapabilityKnowledge::Legacy { thinking: true }
1006        );
1007    }
1008
1009    #[tokio::test]
1010    async fn fallible_discovery_does_not_treat_missing_legacy_support_as_empty_catalog() {
1011        let provider = MockProvider::new("mock");
1012
1013        assert_eq!(
1014            provider.try_discover_models().await.unwrap_err(),
1015            ModelDiscoveryError::Unsupported
1016        );
1017    }
1018
1019    #[test]
1020    fn lifecycle_attach_drops_replaced_providers_after_unlocking_the_owner() {
1021        const PROVIDER_ID: &str = "owner-lock-provider";
1022
1023        struct CatalogCleanup;
1024
1025        impl Drop for CatalogCleanup {
1026            fn drop(&mut self) {
1027                crate::model_registry::remove_provider_catalog(PROVIDER_ID);
1028            }
1029        }
1030
1031        let _registry_lock = crate::model_registry::MODEL_CONFIG_LOCK
1032            .lock()
1033            .unwrap_or_else(std::sync::PoisonError::into_inner);
1034        crate::model_registry::remove_provider_catalog(PROVIDER_ID);
1035        let _catalog_cleanup = CatalogCleanup;
1036        let config = tempfile::tempdir().unwrap();
1037        let hub = crate::config_hub::ConfigHub::from_config_dir(config.path());
1038        let root =
1039            crate::provider_lifecycle::ProviderLifecycle::new(hub.clone(), ProviderRegistry::new());
1040        tokio::runtime::Builder::new_current_thread()
1041            .enable_all()
1042            .build()
1043            .unwrap()
1044            .block_on(root.install_pre_discovered_provider(
1045                crate::auth_store::StoredProvider {
1046                    id: PROVIDER_ID.into(),
1047                    name: "Owner Lock".into(),
1048                    kind: crate::auth_store::ProviderKind::Codex,
1049                    access_token: "access".into(),
1050                    refresh_token: None,
1051                    expires_at: i64::MAX,
1052                    account: None,
1053                    enabled: true,
1054                    model_cache: None,
1055                },
1056                Arc::new(MockProvider::new(PROVIDER_ID)),
1057                vec![DiscoveredModelDetails {
1058                    slug: "owner-lock-model".into(),
1059                    context_budget: Some(128_000),
1060                    capability_knowledge: CapabilityKnowledge::Legacy { thinking: true },
1061                }],
1062            ))
1063            .unwrap();
1064
1065        let target = ProviderRegistry::new();
1066        let owner_was_unlocked = Arc::new(AtomicBool::new(false));
1067        target.register(Arc::new(OwnerLockProbeProvider {
1068            name: PROVIDER_ID.into(),
1069            lifecycle_owner: target.lifecycle_owner.clone(),
1070            owner_was_unlocked: owner_was_unlocked.clone(),
1071        }));
1072
1073        let attached = target.attach_provider_lifecycle(hub).unwrap();
1074        assert!(owner_was_unlocked.load(Ordering::SeqCst));
1075        drop(attached);
1076        root.remove_provider(PROVIDER_ID).unwrap();
1077    }
1078
1079    #[test]
1080    fn resolve_prefix_match_codex_slash_model() {
1081        let reg = fixture_registry();
1082        let p = reg.resolve("codex/gpt-5.6-terra").expect("should resolve");
1083        assert_eq!(p.name(), "codex");
1084    }
1085
1086    #[test]
1087    fn resolve_returns_none_for_unknown() {
1088        let reg = fixture_registry();
1089        assert!(reg.resolve("some-unknown-model").is_none());
1090    }
1091
1092    #[test]
1093    fn remove_updates_membership_and_clears_only_the_removed_default() {
1094        let reg = fixture_registry();
1095        reg.set_default("openai");
1096
1097        assert!(reg.contains("codex"));
1098        assert!(reg.contains("openai"));
1099        assert!(!reg.remove("missing"));
1100        assert_eq!(reg.default.read().unwrap().as_deref(), Some("openai"));
1101
1102        assert!(reg.remove("codex"));
1103        assert!(!reg.contains("codex"));
1104        assert!(reg.contains("openai"));
1105        assert_eq!(reg.default.read().unwrap().as_deref(), Some("openai"));
1106
1107        assert!(reg.remove("openai"));
1108        assert!(!reg.contains("openai"));
1109        assert!(reg.default.read().unwrap().is_none());
1110    }
1111
1112    #[test]
1113    fn resolve_model_registry_provider_field_takes_priority() {
1114        let _registry_lock = crate::model_registry::MODEL_CONFIG_LOCK
1115            .lock()
1116            .unwrap_or_else(std::sync::PoisonError::into_inner);
1117        crate::model_registry::set_provider_config(Default::default());
1118        crate::model_registry::register_model_entries(vec![(
1119            "codex-auto-review".into(),
1120            crate::model_registry::ModelEntry {
1121                model: "codex-auto-review".into(),
1122                provider: Some("codex".into()),
1123                ..Default::default()
1124            },
1125        )]);
1126
1127        let reg = fixture_registry();
1128        let p = reg
1129            .resolve("codex-auto-review")
1130            .expect("should resolve via model registry provider field");
1131        assert_eq!(p.name(), "codex");
1132        crate::model_registry::set_provider_config(Default::default());
1133    }
1134
1135    #[test]
1136    fn resolve_explicit_model_provider_before_slash_prefix() {
1137        let _registry_lock = crate::model_registry::MODEL_CONFIG_LOCK
1138            .lock()
1139            .unwrap_or_else(std::sync::PoisonError::into_inner);
1140        crate::model_registry::set_provider_config(Default::default());
1141        crate::model_registry::register_model_entries(vec![(
1142            "gateway/model".into(),
1143            crate::model_registry::ModelEntry {
1144                model: "api-model".into(),
1145                provider: Some("target".into()),
1146                ..Default::default()
1147            },
1148        )]);
1149
1150        let registry = ProviderRegistry::new();
1151        registry.register(Arc::new(MockProvider::new("gateway")));
1152        registry.register(Arc::new(MockProvider::new("config:target")));
1153
1154        let provider = registry
1155            .resolve("gateway/model")
1156            .expect("explicit model provider should resolve");
1157        assert_eq!(provider.name(), "config:target");
1158
1159        let registry_without_target = ProviderRegistry::new();
1160        registry_without_target.register(Arc::new(MockProvider::new("gateway")));
1161        assert!(registry_without_target.resolve("gateway/model").is_none());
1162        crate::model_registry::set_provider_config(Default::default());
1163    }
1164
1165    #[test]
1166    fn resolve_exact_live_provider_does_not_depend_on_global_auth_store() {
1167        const PROVIDER_ID: &str = "exact-live-selected-config-provider";
1168
1169        struct CatalogCleanup;
1170
1171        impl Drop for CatalogCleanup {
1172            fn drop(&mut self) {
1173                crate::model_registry::remove_provider_catalog(PROVIDER_ID);
1174            }
1175        }
1176
1177        let _registry_lock = crate::model_registry::MODEL_CONFIG_LOCK
1178            .lock()
1179            .unwrap_or_else(std::sync::PoisonError::into_inner);
1180        crate::model_registry::remove_provider_catalog(PROVIDER_ID);
1181        let _catalog_cleanup = CatalogCleanup;
1182        let namespace = "exact-live-selected-config";
1183        let prepared = crate::model_registry::prepare_provider_catalog(
1184            crate::model_registry::ProviderDescriptor {
1185                provider_key: PROVIDER_ID.into(),
1186                provider_name: "Selected config OAuth".into(),
1187                namespace: namespace.into(),
1188                wire_profile: ReasoningWireProfile::CodexResponses,
1189            },
1190            &[DiscoveredModelDetails {
1191                slug: "gpt-selected".into(),
1192                context_budget: Some(128_000),
1193                capability_knowledge: CapabilityKnowledge::Advertised(ModelCapabilities::default()),
1194            }],
1195        )
1196        .unwrap();
1197        crate::model_registry::commit_prepared_provider_catalog(prepared);
1198
1199        let registry = ProviderRegistry::new();
1200        registry.register(Arc::new(MockProvider::new(PROVIDER_ID)));
1201
1202        let provider = registry
1203            .resolve(&format!("{namespace}:gpt-selected"))
1204            .expect("live provider membership should authorize resolution");
1205        assert_eq!(provider.name(), PROVIDER_ID);
1206    }
1207
1208    #[test]
1209    fn reasoning_selection_string_round_trips() {
1210        for value in [
1211            "default",
1212            "off",
1213            "auto",
1214            "auto@pro",
1215            "minimal",
1216            "high@standard",
1217            "xhigh",
1218            "max",
1219            "ultra",
1220            "persistent",
1221            "budget:4096",
1222        ] {
1223            let parsed: ReasoningSelection = value.parse().unwrap();
1224            assert_eq!(parsed.to_string(), value);
1225        }
1226    }
1227
1228    #[test]
1229    fn reasoning_selection_rejects_zero_budget() {
1230        assert!("budget:0".parse::<ReasoningSelection>().is_err());
1231    }
1232
1233    #[test]
1234    fn reasoning_wire_profiles_reject_unrepresentable_controls() {
1235        let high = ReasoningSelection::Effort {
1236            effort: ReasoningEffort::High,
1237            execution_mode: None,
1238        };
1239        assert!(
1240            ReasoningWireProfile::CompatibleThinking
1241                .validate(&high, None)
1242                .unwrap_err()
1243                .contains("cannot represent effort `high`")
1244        );
1245        assert!(
1246            ReasoningWireProfile::CompatibleThinking
1247                .validate(
1248                    &ReasoningSelection::Auto {
1249                        execution_mode: None
1250                    },
1251                    None
1252                )
1253                .is_ok()
1254        );
1255        assert!(
1256            ReasoningWireProfile::AnthropicMessages
1257                .validate(
1258                    &ReasoningSelection::Effort {
1259                        effort: ReasoningEffort::XHigh,
1260                        execution_mode: None,
1261                    },
1262                    None,
1263                )
1264                .unwrap_err()
1265                .contains("cannot represent effort `xhigh`")
1266        );
1267    }
1268}