use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::{Arc, RwLock};
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize, JsonSchema)]
#[serde(rename_all = "kebab-case")]
pub enum ToolProtocol {
#[default]
Native,
TalosStrict,
Compat,
Auto,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ProtocolCapabilities {
pub native_tools: bool,
pub compatibility: bool,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CapabilityProbe {
Known(ProtocolCapabilities),
Unknown,
}
#[derive(Clone, Default)]
pub struct ProtocolCapabilityCache {
entries: Arc<RwLock<HashMap<String, CapabilityProbe>>>,
}
impl ProtocolCapabilityCache {
pub fn new() -> Self {
Self::default()
}
pub fn get(&self, scope: &str) -> Option<CapabilityProbe> {
self.entries.read().ok()?.get(scope).copied()
}
pub fn insert(&self, scope: impl Into<String>, probe: CapabilityProbe) {
if let Ok(mut entries) = self.entries.write() {
entries.insert(scope.into(), probe);
}
}
pub fn clear(&self) {
if let Ok(mut entries) = self.entries.write() {
entries.clear();
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ProtocolFailureDisposition {
Fallback,
Correct,
Stop,
HumanReview,
}
pub fn parse_recovery_decision(input: &str) -> Option<ProtocolFailureDisposition> {
let token = input.trim();
if token.is_empty() || token.contains(char::is_whitespace) {
return None;
}
match token {
"correction" => Some(ProtocolFailureDisposition::Correct),
"fallback" => Some(ProtocolFailureDisposition::Fallback),
"stop" => Some(ProtocolFailureDisposition::Stop),
"human-review" => Some(ProtocolFailureDisposition::HumanReview),
_ => None,
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ExecutionOutcome {
NotStarted,
Completed,
Unknown,
}
impl ExecutionOutcome {
pub const fn permits_retry(self) -> bool {
matches!(self, Self::NotStarted)
}
}
pub fn classify_protocol_failure(
error: &crate::provider::ProviderError,
) -> ProtocolFailureDisposition {
use crate::provider::ProviderError;
match error {
ProviderError::AuthenticationFailed(_)
| ProviderError::RateLimited(_)
| ProviderError::ServerError(_)
| ProviderError::NetworkError(_) => ProtocolFailureDisposition::Stop,
ProviderError::InvalidResponse(_) => ProtocolFailureDisposition::HumanReview,
}
}
impl CapabilityProbe {
pub fn select(self) -> ToolProtocol {
match self {
Self::Known(capabilities) => capabilities.select(),
Self::Unknown => ToolProtocol::Compat,
}
}
}
impl ProtocolCapabilities {
pub fn select(self) -> ToolProtocol {
if self.native_tools {
ToolProtocol::Native
} else if self.compatibility {
ToolProtocol::Compat
} else {
ToolProtocol::TalosStrict
}
}
}
#[cfg(test)]
mod capability_tests {
use super::*;
#[test]
fn selection_prefers_native_then_compat_then_strict() {
assert_eq!(
ProtocolCapabilities {
native_tools: true,
compatibility: true
}
.select(),
ToolProtocol::Native
);
assert_eq!(
ProtocolCapabilities {
native_tools: false,
compatibility: true
}
.select(),
ToolProtocol::Compat
);
assert_eq!(
ProtocolCapabilities {
native_tools: false,
compatibility: false
}
.select(),
ToolProtocol::TalosStrict
);
}
#[test]
fn unknown_probe_uses_compatibility_recovery() {
assert_eq!(CapabilityProbe::Unknown.select(), ToolProtocol::Compat);
}
#[test]
fn failure_classification_is_fail_closed() {
use crate::provider::ProviderError;
for message in [
"invalid API key",
"unsupported authentication protocol",
"malformed",
"timeout",
] {
for error in [
ProviderError::AuthenticationFailed(message.into()),
ProviderError::RateLimited(message.into()),
ProviderError::ServerError(message.into()),
ProviderError::NetworkError(message.into()),
] {
assert_eq!(
classify_protocol_failure(&error),
ProtocolFailureDisposition::Stop
);
}
assert_eq!(
classify_protocol_failure(&ProviderError::InvalidResponse(message.into())),
ProtocolFailureDisposition::HumanReview
);
}
}
#[test]
fn unknown_execution_never_permits_retry() {
assert!(ExecutionOutcome::NotStarted.permits_retry());
assert!(!ExecutionOutcome::Completed.permits_retry());
assert!(!ExecutionOutcome::Unknown.permits_retry());
}
#[test]
fn capability_cache_is_scoped_and_can_cache_unknown() {
let cache = ProtocolCapabilityCache::new();
assert_eq!(cache.get("openai|gpt"), None);
cache.insert("openai|gpt", CapabilityProbe::Unknown);
assert_eq!(cache.get("openai|gpt"), Some(CapabilityProbe::Unknown));
assert_eq!(cache.get("other|gpt"), None);
}
#[test]
fn recovery_decision_requires_exact_token() {
assert_eq!(
parse_recovery_decision("fallback"),
Some(ProtocolFailureDisposition::Fallback)
);
assert_eq!(
parse_recovery_decision("human-review"),
Some(ProtocolFailureDisposition::HumanReview)
);
for value in [
"Fallback",
"fallback now",
"{\"decision\":\"fallback\"}",
"",
] {
assert_eq!(parse_recovery_decision(value), None);
}
}
}
impl ToolProtocol {
pub fn parse(s: &str) -> Option<Self> {
match s {
"native" => Some(ToolProtocol::Native),
"talos-strict" | "talos_xml_json_strict" => Some(ToolProtocol::TalosStrict),
"compat" | "compatibility" => Some(ToolProtocol::Compat),
"auto" => Some(ToolProtocol::Auto),
_ => None,
}
}
}
#[derive(Debug, Clone, Default)]
pub struct ToolProtocolConfig {
pub protocol: ToolProtocol,
pub strict_prompt: bool,
pub stream_filter: bool,
pub schema_validate: bool,
}
impl ToolProtocolConfig {
pub fn for_protocol(protocol: ToolProtocol) -> Self {
match protocol {
ToolProtocol::Native => ToolProtocolConfig {
protocol,
strict_prompt: false,
stream_filter: false,
schema_validate: false,
},
ToolProtocol::TalosStrict => ToolProtocolConfig {
protocol,
strict_prompt: true,
stream_filter: true,
schema_validate: true,
},
ToolProtocol::Compat => ToolProtocolConfig {
protocol,
strict_prompt: false,
stream_filter: true,
schema_validate: false,
},
ToolProtocol::Auto => ToolProtocolConfig {
protocol,
strict_prompt: false,
stream_filter: true,
schema_validate: true,
},
}
}
}