1use std::collections::HashMap;
6use std::sync::atomic::{AtomicU64, Ordering};
7use std::sync::{Arc, RwLock};
8use std::time::{SystemTime, UNIX_EPOCH};
9
10use crate::binary::SignedInvocation;
11use crate::context::DcpContext;
12use crate::dispatch::{BinaryTrieRouter, ServerCapabilities, SharedArgs, ToolResult};
13use crate::security::{NonceStore, SecurityAuditAction, SecurityAuditEvent, SecurityAuditLog};
14use crate::{CapabilityManifest, DCPError, SecurityError};
15
16#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
18pub enum ProtocolVersion {
19 #[default]
21 Mcp,
22 DcpV1,
24}
25
26#[derive(Debug)]
28pub struct Session {
29 pub id: u64,
31 pub protocol: ProtocolVersion,
33 pub created_at: u64,
35 pub last_activity: AtomicU64,
37 pub data: RwLock<HashMap<String, Vec<u8>>>,
39 pub message_count: AtomicU64,
41}
42
43impl Session {
44 pub fn new(id: u64) -> Self {
46 let now = SystemTime::now()
47 .duration_since(UNIX_EPOCH)
48 .unwrap_or_default()
49 .as_secs();
50
51 Self {
52 id,
53 protocol: ProtocolVersion::Mcp,
54 created_at: now,
55 last_activity: AtomicU64::new(now),
56 data: RwLock::new(HashMap::new()),
57 message_count: AtomicU64::new(0),
58 }
59 }
60
61 pub fn touch(&self) {
63 let now = SystemTime::now()
64 .duration_since(UNIX_EPOCH)
65 .unwrap_or_default()
66 .as_secs();
67 self.last_activity.store(now, Ordering::Release);
68 }
69
70 pub fn increment_messages(&self) -> u64 {
72 self.message_count.fetch_add(1, Ordering::AcqRel)
73 }
74
75 pub fn get_data(&self, key: &str) -> Option<Vec<u8>> {
77 self.data.read().ok()?.get(key).cloned()
78 }
79
80 pub fn set_data(&self, key: String, value: Vec<u8>) {
82 if let Ok(mut data) = self.data.write() {
83 data.insert(key, value);
84 }
85 }
86
87 pub fn upgrade_protocol(&mut self, version: ProtocolVersion) {
89 self.protocol = version;
90 }
91
92 pub fn is_dcp(&self) -> bool {
94 matches!(self.protocol, ProtocolVersion::DcpV1)
95 }
96}
97
98#[derive(Debug, Clone)]
100pub struct ServerConfig {
101 pub max_sessions: usize,
103 pub session_timeout_secs: u64,
105 pub enable_metrics: bool,
107 pub server_name: String,
109 pub server_version: String,
111}
112
113impl Default for ServerConfig {
114 fn default() -> Self {
115 Self {
116 max_sessions: 1000,
117 session_timeout_secs: 3600,
118 enable_metrics: true,
119 server_name: "dcp-server".to_string(),
120 server_version: env!("CARGO_PKG_VERSION").to_string(),
121 }
122 }
123}
124
125#[derive(Debug, Default)]
127pub struct Metrics {
128 pub mcp_messages: AtomicU64,
130 pub dcp_messages: AtomicU64,
132 pub mcp_bytes: AtomicU64,
134 pub dcp_bytes: AtomicU64,
136 pub mcp_latency_us: AtomicU64,
138 pub dcp_latency_us: AtomicU64,
140 pub tool_invocations: AtomicU64,
142 pub errors: AtomicU64,
144}
145
146impl Metrics {
147 pub fn record_mcp(&self, bytes: u64, latency_us: u64) {
149 self.mcp_messages.fetch_add(1, Ordering::Relaxed);
150 self.mcp_bytes.fetch_add(bytes, Ordering::Relaxed);
151 self.mcp_latency_us.fetch_add(latency_us, Ordering::Relaxed);
152 }
153
154 pub fn record_dcp(&self, bytes: u64, latency_us: u64) {
156 self.dcp_messages.fetch_add(1, Ordering::Relaxed);
157 self.dcp_bytes.fetch_add(bytes, Ordering::Relaxed);
158 self.dcp_latency_us.fetch_add(latency_us, Ordering::Relaxed);
159 }
160
161 pub fn record_invocation(&self) {
163 self.tool_invocations.fetch_add(1, Ordering::Relaxed);
164 }
165
166 pub fn record_error(&self) {
168 self.errors.fetch_add(1, Ordering::Relaxed);
169 }
170
171 pub fn avg_mcp_latency_us(&self) -> u64 {
173 let count = self.mcp_messages.load(Ordering::Relaxed);
174 if count == 0 {
175 return 0;
176 }
177 self.mcp_latency_us.load(Ordering::Relaxed) / count
178 }
179
180 pub fn avg_dcp_latency_us(&self) -> u64 {
182 let count = self.dcp_messages.load(Ordering::Relaxed);
183 if count == 0 {
184 return 0;
185 }
186 self.dcp_latency_us.load(Ordering::Relaxed) / count
187 }
188
189 pub fn avg_mcp_size(&self) -> u64 {
191 let count = self.mcp_messages.load(Ordering::Relaxed);
192 if count == 0 {
193 return 0;
194 }
195 self.mcp_bytes.load(Ordering::Relaxed) / count
196 }
197
198 pub fn avg_dcp_size(&self) -> u64 {
200 let count = self.dcp_messages.load(Ordering::Relaxed);
201 if count == 0 {
202 return 0;
203 }
204 self.dcp_bytes.load(Ordering::Relaxed) / count
205 }
206
207 pub fn snapshot(&self) -> MetricsSnapshot {
209 MetricsSnapshot {
210 mcp_messages: self.mcp_messages.load(Ordering::Relaxed),
211 dcp_messages: self.dcp_messages.load(Ordering::Relaxed),
212 mcp_bytes: self.mcp_bytes.load(Ordering::Relaxed),
213 dcp_bytes: self.dcp_bytes.load(Ordering::Relaxed),
214 avg_mcp_latency_us: self.avg_mcp_latency_us(),
215 avg_dcp_latency_us: self.avg_dcp_latency_us(),
216 tool_invocations: self.tool_invocations.load(Ordering::Relaxed),
217 errors: self.errors.load(Ordering::Relaxed),
218 }
219 }
220}
221
222#[derive(Debug, Clone)]
224pub struct MetricsSnapshot {
225 pub mcp_messages: u64,
226 pub dcp_messages: u64,
227 pub mcp_bytes: u64,
228 pub dcp_bytes: u64,
229 pub avg_mcp_latency_us: u64,
230 pub avg_dcp_latency_us: u64,
231 pub tool_invocations: u64,
232 pub errors: u64,
233}
234
235pub struct DcpServer {
237 router: BinaryTrieRouter,
239 pub context: Arc<DcpContext>,
241 pub config: ServerConfig,
243 sessions: RwLock<HashMap<u64, Arc<Session>>>,
245 session_counter: AtomicU64,
247 pub metrics: Arc<Metrics>,
249 security_audit: SecurityAuditLog,
251}
252
253impl DcpServer {
254 pub fn new(router: BinaryTrieRouter, context: DcpContext, config: ServerConfig) -> Self {
256 Self {
257 router,
258 context: Arc::new(context),
259 config,
260 sessions: RwLock::new(HashMap::new()),
261 session_counter: AtomicU64::new(1),
262 metrics: Arc::new(Metrics::default()),
263 security_audit: SecurityAuditLog::new(),
264 }
265 }
266
267 pub fn security_audit(&self) -> SecurityAuditLog {
269 self.security_audit.clone()
270 }
271
272 pub fn create_session(&self) -> Result<Arc<Session>, DCPError> {
274 let sessions = self.sessions.read().map_err(|_| DCPError::InternalError)?;
275 if sessions.len() >= self.config.max_sessions {
276 return Err(DCPError::ResourceExhausted);
277 }
278 drop(sessions);
279
280 let id = self.session_counter.fetch_add(1, Ordering::SeqCst);
281 let session = Arc::new(Session::new(id));
282
283 let mut sessions = self.sessions.write().map_err(|_| DCPError::InternalError)?;
284 sessions.insert(id, Arc::clone(&session));
285
286 Ok(session)
287 }
288
289 pub fn get_session(&self, id: u64) -> Option<Arc<Session>> {
291 self.sessions.read().ok()?.get(&id).cloned()
292 }
293
294 pub fn remove_session(&self, id: u64) -> Option<Arc<Session>> {
296 self.sessions.write().ok()?.remove(&id)
297 }
298
299 pub fn session_count(&self) -> usize {
301 self.sessions.read().map(|s| s.len()).unwrap_or(0)
302 }
303
304 pub fn invoke(&self, tool_id: u16, args: &SharedArgs) -> Result<ToolResult, DCPError> {
309 let _ = (tool_id, args);
310 if self.config.enable_metrics {
311 self.metrics.record_error();
312 }
313 self.audit_raw_invoke_denial(tool_id);
314 Err(DCPError::CapabilityDenied)
315 }
316
317 pub fn invoke_authorized(
319 &self,
320 capabilities: &CapabilityManifest,
321 tool_id: u16,
322 args: &SharedArgs,
323 ) -> Result<ToolResult, SecurityError> {
324 if self.config.enable_metrics {
325 self.metrics.record_invocation();
326 }
327
328 self.router.execute_authorized(capabilities, tool_id, args)
329 }
330
331 pub fn invoke_signed_authorized(
334 &self,
335 capabilities: &CapabilityManifest,
336 invocation: &SignedInvocation,
337 public_key: &[u8; 32],
338 nonce_store: &mut NonceStore,
339 args: &SharedArgs,
340 ) -> Result<ToolResult, SecurityError> {
341 let result = self.router.execute_signed_authorized(
342 capabilities,
343 invocation,
344 public_key,
345 nonce_store,
346 args,
347 );
348
349 if self.config.enable_metrics {
350 if result.is_ok() {
351 self.metrics.record_invocation();
352 } else {
353 self.metrics.record_error();
354 }
355 }
356
357 if let Err(error) = result {
358 self.audit_signed_invocation_error(error, invocation);
359 }
360
361 result
362 }
363
364 fn audit_signed_invocation_error(&self, error: SecurityError, invocation: &SignedInvocation) {
365 let (action, reason) = match error {
366 SecurityError::InvalidSignature | SecurityError::ArgsHashMismatch => (
367 SecurityAuditAction::SignatureRejected,
368 match error {
369 SecurityError::ArgsHashMismatch => "args_hash_mismatch",
370 _ => "invalid_signature",
371 },
372 ),
373 SecurityError::ReplayAttack
374 | SecurityError::ExpiredTimestamp
375 | SecurityError::CapacityExceeded => (
376 SecurityAuditAction::ReplayRejected,
377 match error {
378 SecurityError::ReplayAttack => "replay_attack",
379 SecurityError::ExpiredTimestamp => "expired_timestamp",
380 _ => "replay_capacity_exceeded",
381 },
382 ),
383 SecurityError::InsufficientCapabilities => {
384 (SecurityAuditAction::CapabilityDenied, "capability_denied")
385 }
386 SecurityError::ValidationFailed => {
387 (SecurityAuditAction::ValidationRejected, "validation_failed")
388 }
389 };
390
391 self.security_audit.record(
392 SecurityAuditEvent::new(action, reason)
393 .with_method("dcp.tool.invoke_signed_authorized")
394 .with_field("tool_id", invocation.tool_id.to_string()),
395 );
396 }
397
398 pub fn invoke_by_name(&self, name: &str, args: &SharedArgs) -> Result<ToolResult, DCPError> {
403 let _ = (name, args);
404 if self.config.enable_metrics {
405 self.metrics.record_error();
406 }
407 self.audit_raw_invoke_by_name_denial(name);
408 Err(DCPError::CapabilityDenied)
409 }
410
411 fn audit_raw_invoke_denial(&self, tool_id: u16) {
412 self.security_audit.record(
413 SecurityAuditEvent::new(SecurityAuditAction::CapabilityDenied, "raw_invoke_denied")
414 .with_method("dcp.tool.invoke")
415 .with_field("tool_id", tool_id.to_string()),
416 );
417 }
418
419 fn audit_raw_invoke_by_name_denial(&self, name: &str) {
420 self.security_audit.record(
421 SecurityAuditEvent::new(
422 SecurityAuditAction::CapabilityDenied,
423 "raw_invoke_by_name_denied",
424 )
425 .with_method("dcp.tool.invoke_by_name")
426 .with_field("tool_name", name),
427 );
428 }
429
430 pub fn invoke_by_name_authorized(
432 &self,
433 capabilities: &CapabilityManifest,
434 name: &str,
435 args: &SharedArgs,
436 ) -> Result<ToolResult, SecurityError> {
437 let tool_id = self
438 .router
439 .resolve_name(name)
440 .ok_or(SecurityError::InsufficientCapabilities)?;
441
442 self.invoke_authorized(capabilities, tool_id, args)
443 }
444
445 pub fn upgrade_session(&self, session_id: u64) -> Result<(), DCPError> {
447 let _session = self
448 .get_session(session_id)
449 .ok_or(DCPError::SessionNotFound)?;
450
451 let mut sessions = self.sessions.write().map_err(|_| DCPError::InternalError)?;
454 if let Some(session) = sessions.get_mut(&session_id) {
455 let mut new_session = Session::new(session_id);
457 new_session.protocol = ProtocolVersion::DcpV1;
458
459 if let Ok(old_data) = session.data.read() {
461 if let Ok(mut new_data) = new_session.data.write() {
462 for (k, v) in old_data.iter() {
463 new_data.insert(k.clone(), v.clone());
464 }
465 }
466 }
467
468 *session = Arc::new(new_session);
469 }
470
471 Ok(())
472 }
473
474 pub fn cleanup_expired_sessions(&self) -> usize {
476 let now = SystemTime::now()
477 .duration_since(UNIX_EPOCH)
478 .unwrap_or_default()
479 .as_secs();
480
481 let mut sessions = match self.sessions.write() {
482 Ok(s) => s,
483 Err(_) => return 0,
484 };
485
486 let expired: Vec<u64> = sessions
487 .iter()
488 .filter(|(_, session)| {
489 let last = session.last_activity.load(Ordering::Acquire);
490 now - last > self.config.session_timeout_secs
491 })
492 .map(|(id, _)| *id)
493 .collect();
494
495 let count = expired.len();
496 for id in expired {
497 sessions.remove(&id);
498 }
499
500 count
501 }
502
503 pub fn server_info(&self) -> ServerInfo {
505 ServerInfo {
506 name: self.config.server_name.clone(),
507 version: self.config.server_version.clone(),
508 protocol_version: "1.0".to_string(),
509 capabilities: self.router.capabilities(),
510 }
511 }
512}
513
514#[derive(Debug, Clone)]
516pub struct ServerInfo {
517 pub name: String,
518 pub version: String,
519 pub protocol_version: String,
520 pub capabilities: ServerCapabilities,
521}
522
523#[cfg(test)]
524mod tests {
525 use super::*;
526 use crate::dispatch::ToolHandler;
527 use crate::protocol::ToolSchema;
528
529 struct TestHandler;
530
531 impl ToolHandler for TestHandler {
532 fn execute(&self, _args: &SharedArgs) -> Result<ToolResult, DCPError> {
533 Ok(ToolResult::success(vec![1, 2, 3]))
534 }
535
536 fn schema(&self) -> &ToolSchema {
537 static SCHEMA: ToolSchema = ToolSchema {
538 name: "test",
539 id: 1,
540 description: "Test tool",
541 input: crate::protocol::InputSchema {
542 required: 0,
543 fields: Vec::new(),
544 },
545 };
546 &SCHEMA
547 }
548 }
549
550 #[test]
551 fn test_session_creation() {
552 let session = Session::new(1);
553 assert_eq!(session.id, 1);
554 assert_eq!(session.protocol, ProtocolVersion::Mcp);
555 assert!(!session.is_dcp());
556 }
557
558 #[test]
559 fn test_session_data() {
560 let session = Session::new(1);
561 session.set_data("key".to_string(), vec![1, 2, 3]);
562 assert_eq!(session.get_data("key"), Some(vec![1, 2, 3]));
563 assert_eq!(session.get_data("missing"), None);
564 }
565
566 #[test]
567 fn test_session_touch() {
568 let session = Session::new(1);
569 let initial = session.last_activity.load(Ordering::Acquire);
570 std::thread::sleep(std::time::Duration::from_millis(10));
571 session.touch();
572 let updated = session.last_activity.load(Ordering::Acquire);
573 assert!(updated >= initial);
574 }
575
576 #[test]
577 fn test_metrics() {
578 let metrics = Metrics::default();
579
580 metrics.record_mcp(100, 1000);
581 metrics.record_mcp(200, 2000);
582 metrics.record_dcp(50, 500);
583
584 assert_eq!(metrics.mcp_messages.load(Ordering::Relaxed), 2);
585 assert_eq!(metrics.dcp_messages.load(Ordering::Relaxed), 1);
586 assert_eq!(metrics.avg_mcp_latency_us(), 1500);
587 assert_eq!(metrics.avg_dcp_latency_us(), 500);
588 }
589
590 #[test]
591 fn test_server_session_management() {
592 let router = BinaryTrieRouter::new();
593 let context = DcpContext::new(1);
594 let config = ServerConfig {
595 max_sessions: 10,
596 ..Default::default()
597 };
598 let server = DcpServer::new(router, context, config);
599
600 let session = server.create_session().unwrap();
602 assert_eq!(session.id, 1);
603 assert_eq!(server.session_count(), 1);
604
605 let retrieved = server.get_session(1).unwrap();
607 assert_eq!(retrieved.id, 1);
608
609 server.remove_session(1);
611 assert_eq!(server.session_count(), 0);
612 }
613
614 #[test]
615 fn test_server_max_sessions() {
616 let router = BinaryTrieRouter::new();
617 let context = DcpContext::new(1);
618 let config = ServerConfig {
619 max_sessions: 2,
620 ..Default::default()
621 };
622 let server = DcpServer::new(router, context, config);
623
624 server.create_session().unwrap();
625 server.create_session().unwrap();
626
627 let result = server.create_session();
628 assert!(matches!(result, Err(DCPError::ResourceExhausted)));
629 }
630}