use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, RwLock};
use std::time::{SystemTime, UNIX_EPOCH};
use crate::binary::SignedInvocation;
use crate::context::DcpContext;
use crate::dispatch::{BinaryTrieRouter, ServerCapabilities, SharedArgs, ToolResult};
use crate::security::{NonceStore, SecurityAuditAction, SecurityAuditEvent, SecurityAuditLog};
use crate::{CapabilityManifest, DCPError, SecurityError};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum ProtocolVersion {
#[default]
Mcp,
DcpV1,
}
#[derive(Debug)]
pub struct Session {
pub id: u64,
pub protocol: ProtocolVersion,
pub created_at: u64,
pub last_activity: AtomicU64,
pub data: RwLock<HashMap<String, Vec<u8>>>,
pub message_count: AtomicU64,
}
impl Session {
pub fn new(id: u64) -> Self {
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
Self {
id,
protocol: ProtocolVersion::Mcp,
created_at: now,
last_activity: AtomicU64::new(now),
data: RwLock::new(HashMap::new()),
message_count: AtomicU64::new(0),
}
}
pub fn touch(&self) {
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
self.last_activity.store(now, Ordering::Release);
}
pub fn increment_messages(&self) -> u64 {
self.message_count.fetch_add(1, Ordering::AcqRel)
}
pub fn get_data(&self, key: &str) -> Option<Vec<u8>> {
self.data.read().ok()?.get(key).cloned()
}
pub fn set_data(&self, key: String, value: Vec<u8>) {
if let Ok(mut data) = self.data.write() {
data.insert(key, value);
}
}
pub fn upgrade_protocol(&mut self, version: ProtocolVersion) {
self.protocol = version;
}
pub fn is_dcp(&self) -> bool {
matches!(self.protocol, ProtocolVersion::DcpV1)
}
}
#[derive(Debug, Clone)]
pub struct ServerConfig {
pub max_sessions: usize,
pub session_timeout_secs: u64,
pub enable_metrics: bool,
pub server_name: String,
pub server_version: String,
}
impl Default for ServerConfig {
fn default() -> Self {
Self {
max_sessions: 1000,
session_timeout_secs: 3600,
enable_metrics: true,
server_name: "dcp-server".to_string(),
server_version: env!("CARGO_PKG_VERSION").to_string(),
}
}
}
#[derive(Debug, Default)]
pub struct Metrics {
pub mcp_messages: AtomicU64,
pub dcp_messages: AtomicU64,
pub mcp_bytes: AtomicU64,
pub dcp_bytes: AtomicU64,
pub mcp_latency_us: AtomicU64,
pub dcp_latency_us: AtomicU64,
pub tool_invocations: AtomicU64,
pub errors: AtomicU64,
}
impl Metrics {
pub fn record_mcp(&self, bytes: u64, latency_us: u64) {
self.mcp_messages.fetch_add(1, Ordering::Relaxed);
self.mcp_bytes.fetch_add(bytes, Ordering::Relaxed);
self.mcp_latency_us.fetch_add(latency_us, Ordering::Relaxed);
}
pub fn record_dcp(&self, bytes: u64, latency_us: u64) {
self.dcp_messages.fetch_add(1, Ordering::Relaxed);
self.dcp_bytes.fetch_add(bytes, Ordering::Relaxed);
self.dcp_latency_us.fetch_add(latency_us, Ordering::Relaxed);
}
pub fn record_invocation(&self) {
self.tool_invocations.fetch_add(1, Ordering::Relaxed);
}
pub fn record_error(&self) {
self.errors.fetch_add(1, Ordering::Relaxed);
}
pub fn avg_mcp_latency_us(&self) -> u64 {
let count = self.mcp_messages.load(Ordering::Relaxed);
if count == 0 {
return 0;
}
self.mcp_latency_us.load(Ordering::Relaxed) / count
}
pub fn avg_dcp_latency_us(&self) -> u64 {
let count = self.dcp_messages.load(Ordering::Relaxed);
if count == 0 {
return 0;
}
self.dcp_latency_us.load(Ordering::Relaxed) / count
}
pub fn avg_mcp_size(&self) -> u64 {
let count = self.mcp_messages.load(Ordering::Relaxed);
if count == 0 {
return 0;
}
self.mcp_bytes.load(Ordering::Relaxed) / count
}
pub fn avg_dcp_size(&self) -> u64 {
let count = self.dcp_messages.load(Ordering::Relaxed);
if count == 0 {
return 0;
}
self.dcp_bytes.load(Ordering::Relaxed) / count
}
pub fn snapshot(&self) -> MetricsSnapshot {
MetricsSnapshot {
mcp_messages: self.mcp_messages.load(Ordering::Relaxed),
dcp_messages: self.dcp_messages.load(Ordering::Relaxed),
mcp_bytes: self.mcp_bytes.load(Ordering::Relaxed),
dcp_bytes: self.dcp_bytes.load(Ordering::Relaxed),
avg_mcp_latency_us: self.avg_mcp_latency_us(),
avg_dcp_latency_us: self.avg_dcp_latency_us(),
tool_invocations: self.tool_invocations.load(Ordering::Relaxed),
errors: self.errors.load(Ordering::Relaxed),
}
}
}
#[derive(Debug, Clone)]
pub struct MetricsSnapshot {
pub mcp_messages: u64,
pub dcp_messages: u64,
pub mcp_bytes: u64,
pub dcp_bytes: u64,
pub avg_mcp_latency_us: u64,
pub avg_dcp_latency_us: u64,
pub tool_invocations: u64,
pub errors: u64,
}
pub struct DcpServer {
router: BinaryTrieRouter,
pub context: Arc<DcpContext>,
pub config: ServerConfig,
sessions: RwLock<HashMap<u64, Arc<Session>>>,
session_counter: AtomicU64,
pub metrics: Arc<Metrics>,
security_audit: SecurityAuditLog,
}
impl DcpServer {
pub fn new(router: BinaryTrieRouter, context: DcpContext, config: ServerConfig) -> Self {
Self {
router,
context: Arc::new(context),
config,
sessions: RwLock::new(HashMap::new()),
session_counter: AtomicU64::new(1),
metrics: Arc::new(Metrics::default()),
security_audit: SecurityAuditLog::new(),
}
}
pub fn security_audit(&self) -> SecurityAuditLog {
self.security_audit.clone()
}
pub fn create_session(&self) -> Result<Arc<Session>, DCPError> {
let sessions = self.sessions.read().map_err(|_| DCPError::InternalError)?;
if sessions.len() >= self.config.max_sessions {
return Err(DCPError::ResourceExhausted);
}
drop(sessions);
let id = self.session_counter.fetch_add(1, Ordering::SeqCst);
let session = Arc::new(Session::new(id));
let mut sessions = self.sessions.write().map_err(|_| DCPError::InternalError)?;
sessions.insert(id, Arc::clone(&session));
Ok(session)
}
pub fn get_session(&self, id: u64) -> Option<Arc<Session>> {
self.sessions.read().ok()?.get(&id).cloned()
}
pub fn remove_session(&self, id: u64) -> Option<Arc<Session>> {
self.sessions.write().ok()?.remove(&id)
}
pub fn session_count(&self) -> usize {
self.sessions.read().map(|s| s.len()).unwrap_or(0)
}
pub fn invoke(&self, tool_id: u16, args: &SharedArgs) -> Result<ToolResult, DCPError> {
let _ = (tool_id, args);
if self.config.enable_metrics {
self.metrics.record_error();
}
self.audit_raw_invoke_denial(tool_id);
Err(DCPError::CapabilityDenied)
}
pub fn invoke_authorized(
&self,
capabilities: &CapabilityManifest,
tool_id: u16,
args: &SharedArgs,
) -> Result<ToolResult, SecurityError> {
if self.config.enable_metrics {
self.metrics.record_invocation();
}
self.router.execute_authorized(capabilities, tool_id, args)
}
pub fn invoke_signed_authorized(
&self,
capabilities: &CapabilityManifest,
invocation: &SignedInvocation,
public_key: &[u8; 32],
nonce_store: &mut NonceStore,
args: &SharedArgs,
) -> Result<ToolResult, SecurityError> {
let result = self.router.execute_signed_authorized(
capabilities,
invocation,
public_key,
nonce_store,
args,
);
if self.config.enable_metrics {
if result.is_ok() {
self.metrics.record_invocation();
} else {
self.metrics.record_error();
}
}
if let Err(error) = result {
self.audit_signed_invocation_error(error, invocation);
}
result
}
fn audit_signed_invocation_error(&self, error: SecurityError, invocation: &SignedInvocation) {
let (action, reason) = match error {
SecurityError::InvalidSignature | SecurityError::ArgsHashMismatch => (
SecurityAuditAction::SignatureRejected,
match error {
SecurityError::ArgsHashMismatch => "args_hash_mismatch",
_ => "invalid_signature",
},
),
SecurityError::ReplayAttack
| SecurityError::ExpiredTimestamp
| SecurityError::CapacityExceeded => (
SecurityAuditAction::ReplayRejected,
match error {
SecurityError::ReplayAttack => "replay_attack",
SecurityError::ExpiredTimestamp => "expired_timestamp",
_ => "replay_capacity_exceeded",
},
),
SecurityError::InsufficientCapabilities => {
(SecurityAuditAction::CapabilityDenied, "capability_denied")
}
SecurityError::ValidationFailed => {
(SecurityAuditAction::ValidationRejected, "validation_failed")
}
};
self.security_audit.record(
SecurityAuditEvent::new(action, reason)
.with_method("dcp.tool.invoke_signed_authorized")
.with_field("tool_id", invocation.tool_id.to_string()),
);
}
pub fn invoke_by_name(&self, name: &str, args: &SharedArgs) -> Result<ToolResult, DCPError> {
let _ = (name, args);
if self.config.enable_metrics {
self.metrics.record_error();
}
self.audit_raw_invoke_by_name_denial(name);
Err(DCPError::CapabilityDenied)
}
fn audit_raw_invoke_denial(&self, tool_id: u16) {
self.security_audit.record(
SecurityAuditEvent::new(SecurityAuditAction::CapabilityDenied, "raw_invoke_denied")
.with_method("dcp.tool.invoke")
.with_field("tool_id", tool_id.to_string()),
);
}
fn audit_raw_invoke_by_name_denial(&self, name: &str) {
self.security_audit.record(
SecurityAuditEvent::new(
SecurityAuditAction::CapabilityDenied,
"raw_invoke_by_name_denied",
)
.with_method("dcp.tool.invoke_by_name")
.with_field("tool_name", name),
);
}
pub fn invoke_by_name_authorized(
&self,
capabilities: &CapabilityManifest,
name: &str,
args: &SharedArgs,
) -> Result<ToolResult, SecurityError> {
let tool_id = self
.router
.resolve_name(name)
.ok_or(SecurityError::InsufficientCapabilities)?;
self.invoke_authorized(capabilities, tool_id, args)
}
pub fn upgrade_session(&self, session_id: u64) -> Result<(), DCPError> {
let _session = self
.get_session(session_id)
.ok_or(DCPError::SessionNotFound)?;
let mut sessions = self.sessions.write().map_err(|_| DCPError::InternalError)?;
if let Some(session) = sessions.get_mut(&session_id) {
let mut new_session = Session::new(session_id);
new_session.protocol = ProtocolVersion::DcpV1;
if let Ok(old_data) = session.data.read() {
if let Ok(mut new_data) = new_session.data.write() {
for (k, v) in old_data.iter() {
new_data.insert(k.clone(), v.clone());
}
}
}
*session = Arc::new(new_session);
}
Ok(())
}
pub fn cleanup_expired_sessions(&self) -> usize {
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
let mut sessions = match self.sessions.write() {
Ok(s) => s,
Err(_) => return 0,
};
let expired: Vec<u64> = sessions
.iter()
.filter(|(_, session)| {
let last = session.last_activity.load(Ordering::Acquire);
now - last > self.config.session_timeout_secs
})
.map(|(id, _)| *id)
.collect();
let count = expired.len();
for id in expired {
sessions.remove(&id);
}
count
}
pub fn server_info(&self) -> ServerInfo {
ServerInfo {
name: self.config.server_name.clone(),
version: self.config.server_version.clone(),
protocol_version: "1.0".to_string(),
capabilities: self.router.capabilities(),
}
}
}
#[derive(Debug, Clone)]
pub struct ServerInfo {
pub name: String,
pub version: String,
pub protocol_version: String,
pub capabilities: ServerCapabilities,
}
#[cfg(test)]
mod tests {
use super::*;
use crate::dispatch::ToolHandler;
use crate::protocol::ToolSchema;
struct TestHandler;
impl ToolHandler for TestHandler {
fn execute(&self, _args: &SharedArgs) -> Result<ToolResult, DCPError> {
Ok(ToolResult::success(vec![1, 2, 3]))
}
fn schema(&self) -> &ToolSchema {
static SCHEMA: ToolSchema = ToolSchema {
name: "test",
id: 1,
description: "Test tool",
input: crate::protocol::InputSchema {
required: 0,
fields: Vec::new(),
},
};
&SCHEMA
}
}
#[test]
fn test_session_creation() {
let session = Session::new(1);
assert_eq!(session.id, 1);
assert_eq!(session.protocol, ProtocolVersion::Mcp);
assert!(!session.is_dcp());
}
#[test]
fn test_session_data() {
let session = Session::new(1);
session.set_data("key".to_string(), vec![1, 2, 3]);
assert_eq!(session.get_data("key"), Some(vec![1, 2, 3]));
assert_eq!(session.get_data("missing"), None);
}
#[test]
fn test_session_touch() {
let session = Session::new(1);
let initial = session.last_activity.load(Ordering::Acquire);
std::thread::sleep(std::time::Duration::from_millis(10));
session.touch();
let updated = session.last_activity.load(Ordering::Acquire);
assert!(updated >= initial);
}
#[test]
fn test_metrics() {
let metrics = Metrics::default();
metrics.record_mcp(100, 1000);
metrics.record_mcp(200, 2000);
metrics.record_dcp(50, 500);
assert_eq!(metrics.mcp_messages.load(Ordering::Relaxed), 2);
assert_eq!(metrics.dcp_messages.load(Ordering::Relaxed), 1);
assert_eq!(metrics.avg_mcp_latency_us(), 1500);
assert_eq!(metrics.avg_dcp_latency_us(), 500);
}
#[test]
fn test_server_session_management() {
let router = BinaryTrieRouter::new();
let context = DcpContext::new(1);
let config = ServerConfig {
max_sessions: 10,
..Default::default()
};
let server = DcpServer::new(router, context, config);
let session = server.create_session().unwrap();
assert_eq!(session.id, 1);
assert_eq!(server.session_count(), 1);
let retrieved = server.get_session(1).unwrap();
assert_eq!(retrieved.id, 1);
server.remove_session(1);
assert_eq!(server.session_count(), 0);
}
#[test]
fn test_server_max_sessions() {
let router = BinaryTrieRouter::new();
let context = DcpContext::new(1);
let config = ServerConfig {
max_sessions: 2,
..Default::default()
};
let server = DcpServer::new(router, context, config);
server.create_session().unwrap();
server.create_session().unwrap();
let result = server.create_session();
assert!(matches!(result, Err(DCPError::ResourceExhausted)));
}
}