Skip to main content

aurum_core/provider_platform/
conformance.rs

1//! Capability / descriptor conformance hooks (JOE-1936).
2//!
3//! Deterministic checks that a provider's *declared* capabilities do not
4//! contradict its descriptor or (for fakes in tests) its stated network needs.
5//! Full behavioral conformance against live inference is out of scope here.
6
7use 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/// A single conformance failure (honest, actionable).
16#[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
42/// Descriptor network claim must agree with capability `requires_network` /
43/// `local_only_ok` (no lying about offline suitability).
44pub 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
89/// Provider field on capabilities must match the descriptor id.
90pub 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
108/// Operation on capabilities must match the factory direction under test.
109pub 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
126/// Streaming honesty: Aurum must not claim implementation without advertising.
127pub 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
140/// Speaking-rate range honesty when rate is supported.
141pub 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
164/// Run the standard static checks for one descriptor + capabilities pair.
165pub 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
178/// Built-in descriptors must have unique (id, direction) identities.
179pub 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
233/// Run built-in product conformance: unique ids + network claim samples.
234pub 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    /// Fake factory that *lies*: descriptor says LocalOnly, caps require network.
285    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; // LIE relative to LocalOnly descriptor
304            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(&reg).unwrap();
341    }
342
343    #[test]
344    fn unique_identities_on_builtin() {
345        let reg = ProviderRegistry::builtin().unwrap();
346        check_unique_descriptor_identities(&reg).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(&reg).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}