use std::collections::{HashMap, HashSet};
pub type ProtocolVersion = String;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ProtocolMetadata {
pub name: String,
pub version: ProtocolVersion,
pub priority: i32,
pub description: String,
}
impl ProtocolMetadata {
pub fn new(name: impl Into<String>, version: impl Into<String>) -> Self {
Self {
name: name.into(),
version: version.into(),
priority: 0,
description: String::new(),
}
}
pub fn with_priority(mut self, priority: i32) -> Self {
self.priority = priority;
self
}
pub fn with_description(mut self, desc: impl Into<String>) -> Self {
self.description = desc.into();
self
}
pub fn header_value(&self) -> String {
format!("{}.{}", self.name, self.version)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum NegotiationOutcome {
Accepted {
header_value: String,
metadata: ProtocolMetadata,
},
NotRequested,
NoMatch {
requested: Vec<String>,
},
}
impl NegotiationOutcome {
pub fn is_accepted(&self) -> bool {
matches!(self, NegotiationOutcome::Accepted { .. })
}
pub fn header_value(&self) -> Option<&str> {
match self {
NegotiationOutcome::Accepted { header_value, .. } => Some(header_value),
_ => None,
}
}
}
#[derive(Debug, Clone, Default)]
pub struct SubProtocolRegistry {
protocols: Vec<String>,
}
impl SubProtocolRegistry {
pub fn new() -> Self {
Self::default()
}
pub fn register(&mut self, name: &str) {
if !self.protocols.contains(&name.to_string()) {
self.protocols.push(name.to_string());
}
}
pub fn register_many(&mut self, names: &[&str]) {
for name in names {
self.register(name);
}
}
pub fn is_registered(&self, name: &str) -> bool {
self.protocols.contains(&name.to_string())
}
pub fn protocols(&self) -> &[String] {
&self.protocols
}
pub fn negotiate(&self, client_protocols: &[String]) -> Option<String> {
let registered: HashSet<&str> = self.protocols.iter().map(|s| s.as_str()).collect();
client_protocols
.iter()
.find(|p| registered.contains(p.as_str()))
.cloned()
}
pub fn clear(&mut self) {
self.protocols.clear();
}
pub fn len(&self) -> usize {
self.protocols.len()
}
pub fn is_empty(&self) -> bool {
self.protocols.is_empty()
}
}
#[derive(Debug, Clone, Default)]
pub struct NegotiationStats {
pub total_negotiations: u64,
pub accepted: u64,
pub not_requested: u64,
pub no_match: u64,
}
impl NegotiationStats {
pub fn success_rate(&self) -> f64 {
if self.total_negotiations == 0 {
return 0.0;
}
self.accepted as f64 / self.total_negotiations as f64
}
}
pub struct VersionedNegotiator {
protocols: HashMap<String, ProtocolMetadata>,
stats: NegotiationStats,
}
impl Default for VersionedNegotiator {
fn default() -> Self {
Self::new()
}
}
impl VersionedNegotiator {
pub fn new() -> Self {
Self {
protocols: HashMap::new(),
stats: NegotiationStats::default(),
}
}
pub fn register(&mut self, metadata: ProtocolMetadata) {
let key = metadata.header_value();
self.protocols.insert(key, metadata);
}
pub fn register_simple(&mut self, name: &str, version: &str) {
self.register(ProtocolMetadata::new(name, version));
}
pub fn unregister(&mut self, header_value: &str) -> bool {
self.protocols.remove(header_value).is_some()
}
pub fn contains(&self, header_value: &str) -> bool {
self.protocols.contains_key(header_value)
}
pub fn len(&self) -> usize {
self.protocols.len()
}
pub fn is_empty(&self) -> bool {
self.protocols.is_empty()
}
pub fn stats(&self) -> NegotiationStats {
self.stats.clone()
}
pub fn registered_protocols(&self) -> Vec<String> {
let mut keys: Vec<String> = self.protocols.keys().cloned().collect();
keys.sort();
keys
}
pub fn protocols_by_priority(&self) -> Vec<&ProtocolMetadata> {
let mut list: Vec<&ProtocolMetadata> = self.protocols.values().collect();
list.sort_by(|a, b| {
b.priority
.cmp(&a.priority)
.then_with(|| a.header_value().cmp(&b.header_value()))
});
list
}
pub fn negotiate(&mut self, client_protocols: &[String]) -> NegotiationOutcome {
self.stats.total_negotiations += 1;
if client_protocols.is_empty() {
self.stats.not_requested += 1;
return NegotiationOutcome::NotRequested;
}
let candidates: Vec<&ProtocolMetadata> = client_protocols
.iter()
.filter_map(|c| self.protocols.get(c))
.collect();
if candidates.is_empty() {
self.stats.no_match += 1;
return NegotiationOutcome::NoMatch {
requested: client_protocols.to_vec(),
};
}
let best = candidates
.iter()
.max_by_key(|m| m.priority)
.copied()
.expect("candidates is non-empty");
self.stats.accepted += 1;
NegotiationOutcome::Accepted {
header_value: best.header_value(),
metadata: best.clone(),
}
}
pub fn clear(&mut self) {
self.protocols.clear();
self.stats = NegotiationStats::default();
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_subprotocol_registry_new() {
let reg = SubProtocolRegistry::new();
assert!(reg.is_empty());
assert_eq!(reg.len(), 0);
}
#[test]
fn test_subprotocol_register_and_check() {
let mut reg = SubProtocolRegistry::new();
reg.register("json");
assert!(reg.is_registered("json"));
assert!(!reg.is_registered("xml"));
assert_eq!(reg.len(), 1);
}
#[test]
fn test_subprotocol_no_duplicate() {
let mut reg = SubProtocolRegistry::new();
reg.register("json");
reg.register("json");
assert_eq!(reg.len(), 1);
}
#[test]
fn test_subprotocol_register_many() {
let mut reg = SubProtocolRegistry::new();
reg.register_many(&["json", "xml", "protobuf"]);
assert_eq!(reg.len(), 3);
assert!(reg.is_registered("protobuf"));
}
#[test]
fn test_subprotocol_negotiate_matches_first() {
let mut reg = SubProtocolRegistry::new();
reg.register_many(&["json", "protobuf"]);
let client = vec![
"xml".to_string(),
"json".to_string(),
"protobuf".to_string(),
];
let result = reg.negotiate(&client);
assert_eq!(result, Some("json".to_string()));
}
#[test]
fn test_subprotocol_negotiate_no_match() {
let reg = SubProtocolRegistry::new();
let client = vec!["xml".to_string(), "msgpack".to_string()];
let result = reg.negotiate(&client);
assert!(result.is_none());
}
#[test]
fn test_subprotocol_negotiate_empty_client() {
let mut reg = SubProtocolRegistry::new();
reg.register("json");
let client: Vec<String> = vec![];
assert!(reg.negotiate(&client).is_none());
}
#[test]
fn test_subprotocol_negotiate_empty_registry() {
let reg = SubProtocolRegistry::new();
let client = vec!["json".to_string()];
assert!(reg.negotiate(&client).is_none());
}
#[test]
fn test_subprotocol_clear() {
let mut reg = SubProtocolRegistry::new();
reg.register_many(&["json", "xml"]);
assert_eq!(reg.len(), 2);
reg.clear();
assert!(reg.is_empty());
}
#[test]
fn test_subprotocol_protocols_list() {
let mut reg = SubProtocolRegistry::new();
reg.register_many(&["json", "xml"]);
let list = reg.protocols();
assert_eq!(list.len(), 2);
assert!(list.contains(&"json".to_string()));
}
#[test]
fn test_subprotocol_preserves_registration_order() {
let mut reg = SubProtocolRegistry::new();
reg.register("c");
reg.register("a");
reg.register("b");
assert_eq!(reg.protocols(), &["c", "a", "b"]);
}
#[test]
fn test_subprotocol_negotiate_returns_client_order_not_registry_order() {
let mut reg = SubProtocolRegistry::new();
reg.register("a");
reg.register("b");
let client = vec!["b".to_string(), "a".to_string()];
assert_eq!(reg.negotiate(&client), Some("b".to_string()));
}
#[test]
fn test_protocol_metadata_new() {
let meta = ProtocolMetadata::new("chat", "1.0");
assert_eq!(meta.name, "chat");
assert_eq!(meta.version, "1.0");
assert_eq!(meta.priority, 0);
assert!(meta.description.is_empty());
}
#[test]
fn test_protocol_metadata_with_priority() {
let meta = ProtocolMetadata::new("chat", "1.0").with_priority(10);
assert_eq!(meta.priority, 10);
}
#[test]
fn test_protocol_metadata_with_description() {
let meta = ProtocolMetadata::new("chat", "1.0").with_description("Chat protocol v1");
assert_eq!(meta.description, "Chat protocol v1");
}
#[test]
fn test_protocol_metadata_header_value() {
let meta = ProtocolMetadata::new("chat", "1.0");
assert_eq!(meta.header_value(), "chat.1.0");
}
#[test]
fn test_protocol_metadata_header_value_with_complex_version() {
let meta = ProtocolMetadata::new("rpc", "2.1.3");
assert_eq!(meta.header_value(), "rpc.2.1.3");
}
#[test]
fn test_protocol_metadata_builder_chain() {
let meta = ProtocolMetadata::new("jsonrpc", "2.0")
.with_priority(5)
.with_description("JSON-RPC 2.0");
assert_eq!(meta.priority, 5);
assert_eq!(meta.description, "JSON-RPC 2.0");
assert_eq!(meta.header_value(), "jsonrpc.2.0");
}
#[test]
fn test_negotiation_outcome_is_accepted() {
let accepted = NegotiationOutcome::Accepted {
header_value: "chat.1.0".to_string(),
metadata: ProtocolMetadata::new("chat", "1.0"),
};
assert!(accepted.is_accepted());
let not_requested = NegotiationOutcome::NotRequested;
assert!(!not_requested.is_accepted());
let no_match = NegotiationOutcome::NoMatch {
requested: vec!["xml".to_string()],
};
assert!(!no_match.is_accepted());
}
#[test]
fn test_negotiation_outcome_header_value() {
let accepted = NegotiationOutcome::Accepted {
header_value: "chat.1.0".to_string(),
metadata: ProtocolMetadata::new("chat", "1.0"),
};
assert_eq!(accepted.header_value(), Some("chat.1.0"));
let not_requested = NegotiationOutcome::NotRequested;
assert_eq!(not_requested.header_value(), None);
let no_match = NegotiationOutcome::NoMatch { requested: vec![] };
assert_eq!(no_match.header_value(), None);
}
#[test]
fn test_negotiation_stats_default() {
let stats = NegotiationStats::default();
assert_eq!(stats.total_negotiations, 0);
assert_eq!(stats.accepted, 0);
assert_eq!(stats.not_requested, 0);
assert_eq!(stats.no_match, 0);
assert_eq!(stats.success_rate(), 0.0);
}
#[test]
fn test_negotiation_stats_success_rate_all_success() {
let stats = NegotiationStats {
total_negotiations: 10,
accepted: 10,
not_requested: 0,
no_match: 0,
};
assert!((stats.success_rate() - 1.0).abs() < 1e-9);
}
#[test]
fn test_negotiation_stats_success_rate_half() {
let stats = NegotiationStats {
total_negotiations: 10,
accepted: 5,
not_requested: 3,
no_match: 2,
};
assert!((stats.success_rate() - 0.5).abs() < 1e-9);
}
#[test]
fn test_negotiation_stats_success_rate_zero_total() {
let stats = NegotiationStats::default();
assert_eq!(stats.success_rate(), 0.0);
}
#[test]
fn test_versioned_negotiator_new_empty() {
let neg = VersionedNegotiator::new();
assert!(neg.is_empty());
assert_eq!(neg.len(), 0);
let stats = neg.stats();
assert_eq!(stats.total_negotiations, 0);
}
#[test]
fn test_versioned_negotiator_register() {
let mut neg = VersionedNegotiator::new();
neg.register(ProtocolMetadata::new("chat", "1.0"));
assert_eq!(neg.len(), 1);
assert!(neg.contains("chat.1.0"));
}
#[test]
fn test_versioned_negotiator_register_simple() {
let mut neg = VersionedNegotiator::new();
neg.register_simple("jsonrpc", "2.0");
assert!(neg.contains("jsonrpc.2.0"));
assert_eq!(neg.len(), 1);
}
#[test]
fn test_versioned_negotiator_unregister() {
let mut neg = VersionedNegotiator::new();
neg.register_simple("chat", "1.0");
assert!(neg.unregister("chat.1.0"));
assert!(!neg.contains("chat.1.0"));
assert_eq!(neg.len(), 0);
}
#[test]
fn test_versioned_negotiator_unregister_missing() {
let mut neg = VersionedNegotiator::new();
assert!(!neg.unregister("nonexistent"));
}
#[test]
fn test_versioned_negotiator_negotiate_success() {
let mut neg = VersionedNegotiator::new();
neg.register(ProtocolMetadata::new("chat", "1.0").with_priority(5));
let client = vec!["chat.1.0".to_string()];
let outcome = neg.negotiate(&client);
assert!(outcome.is_accepted());
assert_eq!(outcome.header_value(), Some("chat.1.0"));
}
#[test]
fn test_versioned_negotiator_not_requested() {
let mut neg = VersionedNegotiator::new();
neg.register_simple("chat", "1.0");
let client: Vec<String> = vec![];
let outcome = neg.negotiate(&client);
assert_eq!(outcome, NegotiationOutcome::NotRequested);
let stats = neg.stats();
assert_eq!(stats.not_requested, 1);
assert_eq!(stats.total_negotiations, 1);
}
#[test]
fn test_versioned_negotiator_no_match() {
let mut neg = VersionedNegotiator::new();
neg.register_simple("chat", "1.0");
let client = vec!["xml.1.0".to_string(), "msgpack.1.0".to_string()];
let outcome = neg.negotiate(&client);
match outcome {
NegotiationOutcome::NoMatch { requested } => {
assert_eq!(requested, client);
}
_ => panic!("expected NoMatch"),
}
let stats = neg.stats();
assert_eq!(stats.no_match, 1);
}
#[test]
fn test_versioned_negotiator_selects_highest_priority() {
let mut neg = VersionedNegotiator::new();
neg.register(ProtocolMetadata::new("chat", "1.0").with_priority(1));
neg.register(ProtocolMetadata::new("chat", "2.0").with_priority(10));
neg.register(ProtocolMetadata::new("chat", "1.5").with_priority(5));
let client = vec![
"chat.1.0".to_string(),
"chat.2.0".to_string(),
"chat.1.5".to_string(),
];
let outcome = neg.negotiate(&client);
assert_eq!(outcome.header_value(), Some("chat.2.0"));
}
#[test]
fn test_versioned_negotiator_priority_tiebreak_client_order() {
let mut neg = VersionedNegotiator::new();
neg.register(ProtocolMetadata::new("a", "1.0").with_priority(5));
neg.register(ProtocolMetadata::new("b", "1.0").with_priority(5));
let client = vec!["b.1.0".to_string(), "a.1.0".to_string()];
let outcome = neg.negotiate(&client);
let header = outcome.header_value().expect("should accept");
assert!(header == "a.1.0" || header == "b.1.0");
}
#[test]
fn test_versioned_negotiator_stats_tracked_across_calls() {
let mut neg = VersionedNegotiator::new();
neg.register_simple("chat", "1.0");
neg.negotiate(&["chat.1.0".to_string()]);
neg.negotiate(&[]);
neg.negotiate(&["xml.1.0".to_string()]);
neg.negotiate(&["chat.1.0".to_string()]);
let stats = neg.stats();
assert_eq!(stats.total_negotiations, 4);
assert_eq!(stats.accepted, 2);
assert_eq!(stats.not_requested, 1);
assert_eq!(stats.no_match, 1);
assert!((stats.success_rate() - 0.5).abs() < 1e-9);
}
#[test]
fn test_versioned_negotiator_registered_protocols_sorted() {
let mut neg = VersionedNegotiator::new();
neg.register_simple("zebra", "1.0");
neg.register_simple("alpha", "1.0");
neg.register_simple("mango", "1.0");
let list = neg.registered_protocols();
assert_eq!(list, vec!["alpha.1.0", "mango.1.0", "zebra.1.0"]);
}
#[test]
fn test_versioned_negotiator_protocols_by_priority_descending() {
let mut neg = VersionedNegotiator::new();
neg.register(ProtocolMetadata::new("low", "1.0").with_priority(1));
neg.register(ProtocolMetadata::new("high", "1.0").with_priority(10));
neg.register(ProtocolMetadata::new("mid", "1.0").with_priority(5));
let sorted = neg.protocols_by_priority();
assert_eq!(sorted[0].name, "high");
assert_eq!(sorted[1].name, "mid");
assert_eq!(sorted[2].name, "low");
}
#[test]
fn test_versioned_negotiator_protocols_by_priority_tiebreak_alpha() {
let mut neg = VersionedNegotiator::new();
neg.register(ProtocolMetadata::new("zeta", "1.0").with_priority(5));
neg.register(ProtocolMetadata::new("alpha", "1.0").with_priority(5));
let sorted = neg.protocols_by_priority();
assert_eq!(sorted[0].name, "alpha");
assert_eq!(sorted[1].name, "zeta");
}
#[test]
fn test_versioned_negotiator_clear() {
let mut neg = VersionedNegotiator::new();
neg.register_simple("chat", "1.0");
neg.negotiate(&["chat.1.0".to_string()]);
neg.clear();
assert!(neg.is_empty());
let stats = neg.stats();
assert_eq!(stats.total_negotiations, 0);
}
#[test]
fn test_versioned_negotiator_overwrite_registration() {
let mut neg = VersionedNegotiator::new();
neg.register(ProtocolMetadata::new("chat", "1.0").with_priority(1));
neg.register(ProtocolMetadata::new("chat", "1.0").with_priority(10));
assert_eq!(neg.len(), 1);
let client = vec!["chat.1.0".to_string()];
let outcome = neg.negotiate(&client);
if let NegotiationOutcome::Accepted { metadata, .. } = outcome {
assert_eq!(metadata.priority, 10);
} else {
panic!("expected Accepted");
}
}
#[test]
fn test_versioned_negotiator_partial_client_match() {
let mut neg = VersionedNegotiator::new();
neg.register(ProtocolMetadata::new("chat", "1.0").with_priority(5));
neg.register(ProtocolMetadata::new("rpc", "2.0").with_priority(3));
let client = vec![
"xml.1.0".to_string(),
"rpc.2.0".to_string(),
"chat.1.0".to_string(),
];
let outcome = neg.negotiate(&client);
assert_eq!(outcome.header_value(), Some("chat.1.0"));
}
#[test]
fn test_versioned_negotiator_default() {
let neg = VersionedNegotiator::default();
assert!(neg.is_empty());
}
}