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#[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 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 pub input: u64,
450 pub cached_input: u64,
452 pub output: u64,
453 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 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 fn discover_models(&self) -> BoxFut<'static, Vec<DiscoveredModel>> {
559 Box::pin(async { vec![] })
560 }
561
562 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}