Skip to main content

dcp/server/
mod.rs

1//! DCP Server implementation.
2//!
3//! Provides the main server struct with router, context, and session management.
4
5use 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/// Protocol version for DCP
17#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
18pub enum ProtocolVersion {
19    /// MCP JSON-RPC protocol
20    #[default]
21    Mcp,
22    /// DCP binary protocol v1
23    DcpV1,
24}
25
26/// Session state for a connected client
27#[derive(Debug)]
28pub struct Session {
29    /// Unique session ID
30    pub id: u64,
31    /// Current protocol version
32    pub protocol: ProtocolVersion,
33    /// Session creation timestamp
34    pub created_at: u64,
35    /// Last activity timestamp
36    pub last_activity: AtomicU64,
37    /// Custom session data
38    pub data: RwLock<HashMap<String, Vec<u8>>>,
39    /// Message count
40    pub message_count: AtomicU64,
41}
42
43impl Session {
44    /// Create a new session
45    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    /// Update last activity timestamp
62    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    /// Increment message count
71    pub fn increment_messages(&self) -> u64 {
72        self.message_count.fetch_add(1, Ordering::AcqRel)
73    }
74
75    /// Get session data
76    pub fn get_data(&self, key: &str) -> Option<Vec<u8>> {
77        self.data.read().ok()?.get(key).cloned()
78    }
79
80    /// Set session data
81    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    /// Upgrade protocol version
88    pub fn upgrade_protocol(&mut self, version: ProtocolVersion) {
89        self.protocol = version;
90    }
91
92    /// Check if session is using DCP protocol
93    pub fn is_dcp(&self) -> bool {
94        matches!(self.protocol, ProtocolVersion::DcpV1)
95    }
96}
97
98/// Server configuration
99#[derive(Debug, Clone)]
100pub struct ServerConfig {
101    /// Maximum concurrent sessions
102    pub max_sessions: usize,
103    /// Session timeout in seconds
104    pub session_timeout_secs: u64,
105    /// Enable metrics collection
106    pub enable_metrics: bool,
107    /// Server name for identification
108    pub server_name: String,
109    /// Server version
110    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/// Performance metrics for protocol comparison
126#[derive(Debug, Default)]
127pub struct Metrics {
128    /// MCP message count
129    pub mcp_messages: AtomicU64,
130    /// DCP message count
131    pub dcp_messages: AtomicU64,
132    /// MCP total bytes
133    pub mcp_bytes: AtomicU64,
134    /// DCP total bytes
135    pub dcp_bytes: AtomicU64,
136    /// MCP total latency (microseconds)
137    pub mcp_latency_us: AtomicU64,
138    /// DCP total latency (microseconds)
139    pub dcp_latency_us: AtomicU64,
140    /// Tool invocation count
141    pub tool_invocations: AtomicU64,
142    /// Error count
143    pub errors: AtomicU64,
144}
145
146impl Metrics {
147    /// Record an MCP message
148    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    /// Record a DCP message
155    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    /// Record a tool invocation
162    pub fn record_invocation(&self) {
163        self.tool_invocations.fetch_add(1, Ordering::Relaxed);
164    }
165
166    /// Record an error
167    pub fn record_error(&self) {
168        self.errors.fetch_add(1, Ordering::Relaxed);
169    }
170
171    /// Get average MCP latency in microseconds
172    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    /// Get average DCP latency in microseconds
181    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    /// Get average MCP message size
190    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    /// Get average DCP message size
199    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    /// Get snapshot of all metrics
208    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/// Snapshot of metrics at a point in time
223#[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
235/// DCP Server
236pub struct DcpServer {
237    /// Tool router
238    router: BinaryTrieRouter,
239    /// Shared context
240    pub context: Arc<DcpContext>,
241    /// Server configuration
242    pub config: ServerConfig,
243    /// Active sessions
244    sessions: RwLock<HashMap<u64, Arc<Session>>>,
245    /// Session ID counter
246    session_counter: AtomicU64,
247    /// Performance metrics
248    pub metrics: Arc<Metrics>,
249    /// Structured security audit receipts.
250    security_audit: SecurityAuditLog,
251}
252
253impl DcpServer {
254    /// Create a new DCP server
255    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    /// Get structured security audit receipts.
268    pub fn security_audit(&self) -> SecurityAuditLog {
269        self.security_audit.clone()
270    }
271
272    /// Create a new session
273    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    /// Get a session by ID
290    pub fn get_session(&self, id: u64) -> Option<Arc<Session>> {
291        self.sessions.read().ok()?.get(&id).cloned()
292    }
293
294    /// Remove a session
295    pub fn remove_session(&self, id: u64) -> Option<Arc<Session>> {
296        self.sessions.write().ok()?.remove(&id)
297    }
298
299    /// Get active session count
300    pub fn session_count(&self) -> usize {
301        self.sessions.read().map(|s| s.len()).unwrap_or(0)
302    }
303
304    /// Invoke a tool by ID.
305    ///
306    /// Raw invocation is deny-by-default. Use `invoke_authorized` with the
307    /// negotiated capability manifest for execution.
308    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    /// Invoke a tool only when the negotiated capabilities allow it.
318    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    /// Invoke a signed tool call only when signature, args hash, negotiated
332    /// capabilities, schema validation, and replay protection all pass.
333    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    /// Invoke a tool by name (for MCP compatibility).
399    ///
400    /// Raw name-based invocation is deny-by-default. Use
401    /// `invoke_by_name_authorized` after capability negotiation.
402    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    /// Invoke a named tool only when negotiated capabilities allow it.
431    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    /// Upgrade a session from MCP to DCP
446    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        // Session data is preserved during upgrade
452        // Only the protocol version changes
453        let mut sessions = self.sessions.write().map_err(|_| DCPError::InternalError)?;
454        if let Some(session) = sessions.get_mut(&session_id) {
455            // Create new session with upgraded protocol
456            let mut new_session = Session::new(session_id);
457            new_session.protocol = ProtocolVersion::DcpV1;
458
459            // Copy over session data
460            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    /// Clean up expired sessions
475    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    /// Get server info
504    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/// Server information for capability negotiation
515#[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        // Create session
601        let session = server.create_session().unwrap();
602        assert_eq!(session.id, 1);
603        assert_eq!(server.session_count(), 1);
604
605        // Get session
606        let retrieved = server.get_session(1).unwrap();
607        assert_eq!(retrieved.id, 1);
608
609        // Remove session
610        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}