aurum_core/provider_platform/
conformance.rs1use super::descriptor::{NetworkRequirement, ProviderDescriptor};
8use super::id::ProviderId;
9use super::registry::ProviderRegistry;
10use crate::capabilities::{CapabilityOperation, ProviderCapabilities};
11use crate::error::{Result, UserError};
12use std::collections::HashSet;
13use std::fmt;
14
15#[derive(Debug, Clone, PartialEq, Eq)]
17pub struct ConformanceFailure {
18 pub provider: String,
19 pub check: &'static str,
20 pub detail: String,
21}
22
23impl fmt::Display for ConformanceFailure {
24 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
25 write!(
26 f,
27 "capability conformance failed for '{}': {} — {}",
28 self.provider, self.check, self.detail
29 )
30 }
31}
32
33impl From<ConformanceFailure> for crate::error::TranscriptionError {
34 fn from(c: ConformanceFailure) -> Self {
35 UserError::Other {
36 message: c.to_string(),
37 }
38 .into()
39 }
40}
41
42pub fn check_network_claim(
45 desc: &ProviderDescriptor,
46 caps: &ProviderCapabilities,
47) -> std::result::Result<(), ConformanceFailure> {
48 match desc.network {
49 NetworkRequirement::LocalOnly => {
50 if caps.requires_network {
51 return Err(ConformanceFailure {
52 provider: desc.id.as_str().into(),
53 check: "network_claim",
54 detail: "descriptor is LocalOnly but capabilities.requires_network is true"
55 .into(),
56 });
57 }
58 if !caps.local_only_ok {
59 return Err(ConformanceFailure {
60 provider: desc.id.as_str().into(),
61 check: "network_claim",
62 detail: "descriptor is LocalOnly but capabilities.local_only_ok is false"
63 .into(),
64 });
65 }
66 }
67 NetworkRequirement::RequiresNetwork => {
68 if !caps.requires_network {
69 return Err(ConformanceFailure {
70 provider: desc.id.as_str().into(),
71 check: "network_claim",
72 detail: "descriptor RequiresNetwork but capabilities.requires_network is false"
73 .into(),
74 });
75 }
76 if caps.local_only_ok {
77 return Err(ConformanceFailure {
78 provider: desc.id.as_str().into(),
79 check: "network_claim",
80 detail: "descriptor RequiresNetwork but capabilities.local_only_ok is true"
81 .into(),
82 });
83 }
84 }
85 }
86 Ok(())
87}
88
89pub fn check_provider_identity(
91 desc: &ProviderDescriptor,
92 caps: &ProviderCapabilities,
93) -> std::result::Result<(), ConformanceFailure> {
94 if caps.provider != desc.id.as_str() {
95 return Err(ConformanceFailure {
96 provider: desc.id.as_str().into(),
97 check: "provider_identity",
98 detail: format!(
99 "capabilities.provider '{}' != descriptor id '{}'",
100 caps.provider,
101 desc.id.as_str()
102 ),
103 });
104 }
105 Ok(())
106}
107
108pub fn check_operation(
110 expected: CapabilityOperation,
111 caps: &ProviderCapabilities,
112) -> std::result::Result<(), ConformanceFailure> {
113 if caps.operation != expected {
114 return Err(ConformanceFailure {
115 provider: caps.provider.clone(),
116 check: "operation",
117 detail: format!(
118 "expected operation {:?}, got {:?}",
119 expected, caps.operation
120 ),
121 });
122 }
123 Ok(())
124}
125
126pub fn check_streaming_honesty(
128 caps: &ProviderCapabilities,
129) -> std::result::Result<(), ConformanceFailure> {
130 if caps.streaming_implemented_by_aurum && !caps.streaming_advertised {
131 return Err(ConformanceFailure {
132 provider: caps.provider.clone(),
133 check: "streaming_honesty",
134 detail: "streaming_implemented_by_aurum requires streaming_advertised".into(),
135 });
136 }
137 Ok(())
138}
139
140pub fn check_speaking_rate_range(
142 caps: &ProviderCapabilities,
143) -> std::result::Result<(), ConformanceFailure> {
144 if caps.supports_speaking_rate {
145 match (caps.speaking_rate_min, caps.speaking_rate_max) {
146 (Some(min), Some(max))
147 if min > 0.0 && max >= min && f32::is_finite(max) && f32::is_finite(min) =>
148 {
149 Ok(())
150 }
151 _ => Err(ConformanceFailure {
152 provider: caps.provider.clone(),
153 check: "speaking_rate_range",
154 detail:
155 "supports_speaking_rate requires finite speaking_rate_min/max with min<=max"
156 .into(),
157 }),
158 }
159 } else {
160 Ok(())
161 }
162}
163
164pub fn check_descriptor_capabilities(
166 desc: &ProviderDescriptor,
167 caps: &ProviderCapabilities,
168 expected_op: CapabilityOperation,
169) -> std::result::Result<(), ConformanceFailure> {
170 check_provider_identity(desc, caps)?;
171 check_operation(expected_op, caps)?;
172 check_network_claim(desc, caps)?;
173 check_streaming_honesty(caps)?;
174 check_speaking_rate_range(caps)?;
175 Ok(())
176}
177
178pub fn check_unique_descriptor_identities(registry: &ProviderRegistry) -> Result<()> {
180 let mut stt_ids: HashSet<String> = HashSet::new();
181 for id in registry.list_stt_ids() {
182 if !stt_ids.insert(id.as_str().to_string()) {
183 return Err(ConformanceFailure {
184 provider: id.as_str().into(),
185 check: "unique_stt_id",
186 detail: format!("duplicate STT provider id '{id}'"),
187 }
188 .into());
189 }
190 }
191
192 #[cfg(feature = "tts")]
193 {
194 let mut tts_ids: HashSet<String> = HashSet::new();
195 for id in registry.list_tts_ids() {
196 if !tts_ids.insert(id.as_str().to_string()) {
197 return Err(ConformanceFailure {
198 provider: id.as_str().into(),
199 check: "unique_tts_id",
200 detail: format!("duplicate TTS provider id '{id}'"),
201 }
202 .into());
203 }
204 }
205 }
206
207 let mut seen: HashSet<String> = HashSet::new();
208 for d in registry.descriptors() {
209 if !seen.insert(d.id.as_str().to_string()) {
210 return Err(ConformanceFailure {
211 provider: d.id.as_str().into(),
212 check: "unique_descriptor_enumeration",
213 detail: format!("duplicate descriptor enumeration for '{}'", d.id),
214 }
215 .into());
216 }
217 }
218
219 for id in &stt_ids {
220 if !seen.contains(id) {
221 return Err(ConformanceFailure {
222 provider: id.clone(),
223 check: "descriptor_coverage",
224 detail: format!("STT id '{id}' missing from descriptors()"),
225 }
226 .into());
227 }
228 }
229
230 Ok(())
231}
232
233pub fn check_builtin_conformance(registry: &ProviderRegistry) -> Result<()> {
235 check_unique_descriptor_identities(registry)?;
236
237 for id in registry.list_stt_ids() {
238 let factory = registry.stt_factory(&id)?;
239 let desc = factory.descriptor();
240 let model = sample_model_for(&id, CapabilityOperation::Stt);
241 let caps = factory.capabilities(model)?;
242 check_descriptor_capabilities(desc, &caps, CapabilityOperation::Stt)?;
243 }
244
245 #[cfg(feature = "tts")]
246 {
247 for id in registry.list_tts_ids() {
248 let factory = registry.tts_factory(&id)?;
249 let desc = factory.descriptor();
250 let model = sample_model_for(&id, CapabilityOperation::Tts);
251 let caps = factory.capabilities(model)?;
252 check_descriptor_capabilities(desc, &caps, CapabilityOperation::Tts)?;
253 }
254 }
255
256 Ok(())
257}
258
259fn sample_model_for(id: &ProviderId, op: CapabilityOperation) -> &'static str {
260 match (id.as_str(), op) {
261 ("local", CapabilityOperation::Stt) => "tiny-q5_1",
262 ("local", CapabilityOperation::Tts) => "kitten-nano-int8",
263 ("openrouter", CapabilityOperation::Stt) => "openai/whisper-large-v3",
264 ("openrouter", CapabilityOperation::Tts) => "openai/gpt-4o-mini-tts-2025-12-15",
265 ("openai", CapabilityOperation::Stt) => "whisper-1",
266 ("openai", CapabilityOperation::Tts) => "tts-1",
267 ("elevenlabs", CapabilityOperation::Tts) => "eleven_multilingual_v2",
268 ("xai", CapabilityOperation::Stt) => "xai-stt",
269 ("xai", CapabilityOperation::Tts) => "xai-tts",
270 _ => "default",
271 }
272}
273
274#[cfg(test)]
275mod tests {
276 use super::super::descriptor::{ProviderOperations, ProviderStability};
277 use super::super::factory::TranscriptionProviderFactory;
278 use super::super::registry::ProviderRegistryBuilder;
279 use super::*;
280 use crate::capabilities::{DescriptorFreshness, SttBackendClass};
281 use crate::providers::TranscriptionProvider;
282 use std::sync::Arc;
283
284 struct LyingNetworkStt {
286 desc: ProviderDescriptor,
287 }
288
289 impl TranscriptionProviderFactory for LyingNetworkStt {
290 fn descriptor(&self) -> &ProviderDescriptor {
291 &self.desc
292 }
293
294 fn capabilities(&self, model: &str) -> Result<ProviderCapabilities> {
295 let mut caps = ProviderCapabilities::with_core(
296 self.desc.id.as_str(),
297 model,
298 CapabilityOperation::Stt,
299 );
300 caps.stt_backend = Some(SttBackendClass::Asr);
301 caps.timestamps_reliable = true;
302 caps.languages = vec!["en".into()];
303 caps.requires_network = true; caps.local_only_ok = false;
305 caps.output_formats = vec!["txt".into()];
306 caps.descriptor_freshness = DescriptorFreshness::Static;
307 Ok(caps)
308 }
309
310 fn build(
311 &self,
312 _ctx: &super::super::context::ProviderBuildContext,
313 ) -> Result<Arc<dyn TranscriptionProvider>> {
314 Err(UserError::Other {
315 message: "lying fake has no inference".into(),
316 }
317 .into())
318 }
319 }
320
321 #[test]
322 fn lying_requires_network_fails_conformance() {
323 let desc = ProviderDescriptor::new(
324 ProviderId::must("liar"),
325 "Liar",
326 ProviderOperations::STT_ONLY,
327 NetworkRequirement::LocalOnly,
328 ProviderStability::TestOnly,
329 );
330 let factory = LyingNetworkStt { desc };
331 let caps = factory.capabilities("x").unwrap();
332 let err = check_network_claim(factory.descriptor(), &caps).unwrap_err();
333 assert_eq!(err.check, "network_claim");
334 assert!(err.detail.contains("requires_network"));
335 }
336
337 #[test]
338 fn builtin_passes_conformance() {
339 let reg = ProviderRegistry::builtin().unwrap();
340 check_builtin_conformance(®).unwrap();
341 }
342
343 #[test]
344 fn unique_identities_on_builtin() {
345 let reg = ProviderRegistry::builtin().unwrap();
346 check_unique_descriptor_identities(®).unwrap();
347 let ids: Vec<_> = reg.descriptors().iter().map(|d| d.id.as_str()).collect();
348 assert!(ids.contains(&"local"));
349 assert!(ids.contains(&"openrouter"));
350 let set: HashSet<_> = ids.iter().copied().collect();
351 assert_eq!(set.len(), ids.len());
352 }
353
354 #[test]
355 fn lying_factory_fails_when_registered_and_checked() {
356 let factory: Arc<dyn TranscriptionProviderFactory> = Arc::new(LyingNetworkStt {
357 desc: ProviderDescriptor::new(
358 ProviderId::must("liar"),
359 "Liar",
360 ProviderOperations::STT_ONLY,
361 NetworkRequirement::LocalOnly,
362 ProviderStability::TestOnly,
363 ),
364 });
365 let reg = ProviderRegistryBuilder::default()
366 .register_stt(factory)
367 .unwrap()
368 .build();
369 let err = check_builtin_conformance(®).unwrap_err();
370 let msg = err.to_string();
371 assert!(
372 msg.contains("network") || msg.contains("conformance"),
373 "{msg}"
374 );
375 }
376
377 #[test]
378 fn streaming_honesty_rejects_implemented_without_advertised() {
379 let mut caps = ProviderCapabilities::with_core("x", "m", CapabilityOperation::Tts);
380 caps.streaming_implemented_by_aurum = true;
381 caps.streaming_advertised = false;
382 let err = check_streaming_honesty(&caps).unwrap_err();
383 assert_eq!(err.check, "streaming_honesty");
384 }
385}