Skip to main content

aurum_core/provider_platform/
mod.rs

1//! Shared provider platform: identity, registry, factories (JOE-1933 / JOE-1932 / JOE-1936).
2//!
3//! Direction-specific execution traits remain in [`crate::providers`] (STT) and
4//! [`crate::tts::provider`] (TTS). This module owns **how** those implementations
5//! are identified, registered, and constructed without flattening STT/TTS
6//! semantics. See `docs/development/adr-002-provider-registry.md`.
7//!
8//! Capability discovery and conformance (JOE-1936) route through the registry
9//! via [`capabilities_for`] and [`conformance`].
10
11mod builtin;
12mod conformance;
13mod context;
14mod descriptor;
15mod factory;
16mod id;
17mod listing;
18mod lookup;
19mod qualification;
20mod registry;
21
22pub use builtin::{LocalSttFactory, OpenRouterSttFactory};
23#[cfg(feature = "tts")]
24pub use builtin::{LocalTtsFactory, OpenRouterTtsFactory};
25pub use conformance::{
26    check_builtin_conformance, check_descriptor_capabilities, check_network_claim,
27    check_unique_descriptor_identities, ConformanceFailure,
28};
29pub use context::ProviderBuildContext;
30pub use descriptor::{
31    NetworkRequirement, ProviderDescriptor, ProviderOperations, ProviderStability,
32};
33
34/// Optional overrides when resolving a provider through [`crate::AurumEngine`] (JOE-1938).
35#[derive(Debug, Clone, Default)]
36pub struct ProviderResolveOptions {
37    /// When true, local STT/TTS may emit progress on stderr.
38    pub show_progress: bool,
39    /// Override OpenRouter STT mode; default parses config.
40    pub stt_mode: Option<crate::providers::OpenRouterSttMode>,
41    /// Override `local_only`; default uses config.
42    pub local_only: Option<bool>,
43}
44#[cfg(feature = "tts")]
45pub use factory::SynthesisProviderFactory;
46pub use factory::TranscriptionProviderFactory;
47pub use id::{ProviderId, MAX_PROVIDER_ID_LEN};
48pub use listing::{
49    list_provider_summaries, merge_provider_summaries, provider_list, ProviderList,
50    ProviderSummary, PROVIDER_LIST_SCHEMA_VERSION,
51};
52pub use lookup::{capabilities_for, preflight_stt_with_registry, preflight_tts_with_registry};
53pub use registry::{ProviderRegistry, ProviderRegistryBuilder};
54
55#[cfg(test)]
56mod tests {
57    use super::*;
58    use crate::capabilities::{CapabilityOperation, ProviderCapabilities};
59    use crate::secret::SecretString;
60    use std::sync::Arc;
61
62    struct FakeStt {
63        desc: ProviderDescriptor,
64    }
65
66    impl TranscriptionProviderFactory for FakeStt {
67        fn descriptor(&self) -> &ProviderDescriptor {
68            &self.desc
69        }
70
71        fn capabilities(&self, model: &str) -> crate::error::Result<ProviderCapabilities> {
72            let mut caps = ProviderCapabilities::with_core(
73                self.desc.id.as_str(),
74                model,
75                CapabilityOperation::Stt,
76            );
77            caps.stt_backend = Some(crate::capabilities::SttBackendClass::Asr);
78            caps.timestamps_reliable = true;
79            caps.languages = vec!["en".into()];
80            caps.max_duration_secs = Some(60.0);
81            caps.supports_cancellation = true;
82            caps.requires_network = false;
83            caps.local_only_ok = true;
84            caps.output_formats = vec!["txt".into()];
85            Ok(caps)
86        }
87
88        fn build(
89            &self,
90            _ctx: &ProviderBuildContext,
91        ) -> crate::error::Result<Arc<dyn crate::providers::TranscriptionProvider>> {
92            Err(crate::error::UserError::Other {
93                message: "fake stt does not implement inference".into(),
94            }
95            .into())
96        }
97    }
98
99    #[test]
100    fn rejects_duplicate_stt_registration() {
101        let f1: Arc<dyn TranscriptionProviderFactory> = Arc::new(FakeStt {
102            desc: ProviderDescriptor::new(
103                ProviderId::must("fake"),
104                "Fake",
105                ProviderOperations::STT_ONLY,
106                NetworkRequirement::LocalOnly,
107                ProviderStability::TestOnly,
108            ),
109        });
110        let f2: Arc<dyn TranscriptionProviderFactory> = Arc::new(FakeStt {
111            desc: ProviderDescriptor::new(
112                ProviderId::must("fake"),
113                "Fake2",
114                ProviderOperations::STT_ONLY,
115                NetworkRequirement::LocalOnly,
116                ProviderStability::TestOnly,
117            ),
118        });
119        let err = match ProviderRegistry::builder()
120            .register_stt(f1)
121            .unwrap()
122            .register_stt(f2)
123        {
124            Ok(_) => panic!("expected duplicate registration error"),
125            Err(e) => e,
126        };
127        let msg = err.to_string();
128        assert!(msg.contains("duplicate"), "{msg}");
129    }
130
131    #[test]
132    fn builtin_registers_local_and_openrouter_stt() {
133        let reg = ProviderRegistry::builtin().unwrap();
134        let local = ProviderId::local();
135        let or = ProviderId::openrouter();
136        assert!(reg.stt_factory(&local).is_ok());
137        assert!(reg.stt_factory(&or).is_ok());
138        let caps = reg
139            .stt_factory(&local)
140            .unwrap()
141            .capabilities("tiny-q5_1")
142            .unwrap();
143        assert_eq!(caps.provider, "local");
144        assert!(!caps.requires_network || caps.local_only_ok);
145
146        assert!(reg.stt_factory(&ProviderId::must("openai")).is_ok());
147        assert!(reg.stt_factory(&ProviderId::must("elevenlabs")).is_err());
148    }
149
150    #[test]
151    #[cfg(feature = "tts")]
152    fn builtin_registers_local_and_openrouter_tts() {
153        let reg = ProviderRegistry::builtin().unwrap();
154        let local = ProviderId::local();
155        assert!(reg.tts_factory(&local).is_ok());
156        // OpenRouter remote TTS (JOE-1939) + first-party OpenAI (JOE-1940).
157        assert!(reg.tts_factory(&ProviderId::openrouter()).is_ok());
158        assert!(reg.tts_factory(&ProviderId::must("openai")).is_ok());
159        assert!(reg.stt_factory(&ProviderId::must("openai")).is_ok());
160    }
161
162    #[test]
163    fn descriptors_are_deterministic() {
164        let a = ProviderRegistry::builtin().unwrap();
165        let b = ProviderRegistry::builtin().unwrap();
166        let da: Vec<_> = a.descriptors().iter().map(|d| d.id.as_str()).collect();
167        let db: Vec<_> = b.descriptors().iter().map(|d| d.id.as_str()).collect();
168        assert_eq!(da, db);
169        assert!(da.contains(&"local"));
170        assert!(da.contains(&"openrouter"));
171    }
172
173    #[test]
174    fn openrouter_factory_rejects_local_only() {
175        let reg = ProviderRegistry::builtin().unwrap();
176        let f = reg.stt_factory(&ProviderId::openrouter()).unwrap();
177        let ctx = ProviderBuildContext::new("/tmp/aurum-reg-test")
178            .with_local_only(true)
179            .with_api_key(Some(SecretString::new("sk-test")));
180        let err = match f.build(&ctx) {
181            Ok(_) => panic!("expected local_only rejection"),
182            Err(e) => e,
183        };
184        let msg = err.to_string();
185        assert!(
186            msg.contains("local_only") || msg.contains("remote"),
187            "{msg}"
188        );
189    }
190
191    #[test]
192    fn local_stt_builds_without_secret() {
193        let reg = ProviderRegistry::builtin().unwrap();
194        let f = reg.stt_factory(&ProviderId::local()).unwrap();
195        let ctx = ProviderBuildContext::new(std::env::temp_dir().join("aurum-reg-local"));
196        let p = f.build(&ctx).unwrap();
197        assert_eq!(p.name(), "local");
198    }
199
200    #[test]
201    fn secret_scoping_does_not_leak_via_debug() {
202        let ctx = ProviderBuildContext::new("/tmp/x")
203            .with_api_key(Some(SecretString::new("sk-must-not-appear-in-debug")));
204        assert!(!format!("{ctx:?}").contains("sk-must-not-appear"));
205    }
206
207    #[test]
208    fn capabilities_for_routes_through_registry() {
209        let reg = ProviderRegistry::builtin().unwrap();
210        let caps =
211            capabilities_for(&reg, &ProviderId::local(), CapabilityOperation::Stt, "base").unwrap();
212        assert_eq!(caps.provider, "local");
213        assert_eq!(
214            caps.schema_version,
215            crate::capabilities::CAPABILITY_SCHEMA_VERSION
216        );
217    }
218}