1use crate::error::{Result, WebmcpError};
2use crate::event_hub::{EventHubConfig, EventHubSubscription, MAX_EVENT_BYTES, WebmcpEventHub};
3use crate::pairing::{PairingDisplay, PairingManager, is_valid_origin};
4use crate::protocol::{
5 BridgeEventMessage, BridgeRequest, BridgeResponse, BridgeSettings, PROTOCOL_VERSION, PairPayload, StatusPayload,
6 is_valid_request_id, response_request_id,
7};
8use crate::remote_mcp::{RemoteMcpEndpoint, RemoteMcpServerConfig};
9use crate::runtime::RuntimeAdapter;
10use axum::Router;
11use axum::extract::{State, WebSocketUpgrade, ws};
12use axum::http::{HeaderMap, StatusCode, header::ORIGIN};
13use axum::response::{IntoResponse, Response};
14use futures::{StreamExt, future::BoxFuture};
15use serde_json::Value;
16use std::net::IpAddr;
17use std::sync::Arc;
18use std::time::{Duration, Instant};
19use tokio::net::TcpListener;
20use tokio::sync::{OwnedSemaphorePermit, Semaphore, mpsc, oneshot};
21
22const MAX_PAIRED_CONNECTIONS: usize = 64;
23const MAX_CONNECTIONS: usize = 128;
24const HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(30);
25const EVENT_ENVELOPE_OVERHEAD: usize = 256;
26const MIN_FRAME_BYTES: usize = 1024;
27const WRITE_TIMEOUT: Duration = Duration::from_secs(10);
28const MAX_MUTATION_QUEUE: usize = 64;
29const SESSION_LEASE_CHECK_INTERVAL: Duration = Duration::from_millis(250);
30
31#[derive(Debug, Clone)]
33pub struct WebmcpServerConfig {
34 pub host: String,
36 pub port: u16,
38 pub allowed_origins: Vec<String>,
40 pub pairing_ttl_secs: u64,
42 pub max_frame_bytes: usize,
44 pub max_in_flight_requests: usize,
46 pub allow_remote: bool,
48 pub public_url: Option<String>,
50 pub remote_mcp: Option<RemoteMcpServerConfig>,
52 pub event_hub: EventHubConfig,
54 pub request_timeout: Duration,
56}
57
58impl Default for WebmcpServerConfig {
59 fn default() -> Self {
60 Self {
61 host: "127.0.0.1".to_string(),
62 port: 0,
63 allowed_origins: Vec::new(),
64 pairing_ttl_secs: 300,
65 max_frame_bytes: 1024 * 1024,
66 max_in_flight_requests: 8,
67 allow_remote: false,
68 public_url: None,
69 remote_mcp: None,
70 event_hub: EventHubConfig::default(),
71 request_timeout: Duration::from_secs(30),
72 }
73 }
74}
75
76struct ServerState {
77 adapter: Arc<dyn RuntimeAdapter>,
78 pairing: PairingManager,
79 event_hub: WebmcpEventHub,
80 dispatch: Arc<DispatchState>,
81 mutation_supervisor: MutationSupervisor,
82 paired_connections: Arc<Semaphore>,
83 connections: Arc<Semaphore>,
84 max_frame_bytes: usize,
85 request_timeout: Duration,
86 remote_mcp: Option<Arc<RemoteMcpEndpoint>>,
87}
88
89struct DispatchState {
90 adapter: Arc<dyn RuntimeAdapter>,
91 pairing: PairingManager,
92 event_hub: WebmcpEventHub,
93 settings: BridgeSettings,
94 in_flight: Arc<Semaphore>,
95 request_timeout: Duration,
96}
97
98struct MutationJob {
99 dispatch: Arc<DispatchState>,
100 origin: String,
101 token: String,
102 request: BridgeRequest,
103 result: oneshot::Sender<Result<Value>>,
104}
105
106#[derive(Clone, Default)]
107struct MutationSupervisor {
108 sender: Arc<tokio::sync::Mutex<Option<mpsc::Sender<MutationJob>>>>,
109}
110
111impl MutationSupervisor {
112 async fn submit(
113 &self,
114 dispatch: Arc<DispatchState>,
115 origin: String,
116 token: String,
117 request: BridgeRequest,
118 ) -> Result<Value> {
119 let sender = {
120 let mut sender_slot = self.sender.lock().await;
121 if let Some(sender) = sender_slot.as_ref() {
122 sender.clone()
123 } else {
124 let (sender, receiver) = mpsc::channel(MAX_MUTATION_QUEUE);
125 drop(tokio::spawn(run_mutation_supervisor(receiver)));
126 *sender_slot = Some(sender.clone());
127 sender
128 }
129 };
130 let (result, receiver) = oneshot::channel();
131 sender
132 .try_send(MutationJob { dispatch, origin, token, request, result })
133 .map_err(|error| match error {
134 mpsc::error::TrySendError::Full(_) => WebmcpError::LimitExceeded,
135 mpsc::error::TrySendError::Closed(_) => {
136 WebmcpError::Adapter("WebMCP mutation supervisor is closed".to_string())
137 }
138 })?;
139 receiver
140 .await
141 .map_err(|_error| WebmcpError::Adapter("WebMCP mutation supervisor stopped".to_string()))?
142 }
143}
144
145async fn run_mutation_supervisor(mut receiver: mpsc::Receiver<MutationJob>) {
146 while let Some(job) = receiver.recv().await {
147 let result = dispatch_request(job.dispatch, job.origin, job.token, job.request).await;
148 drop(job.result.send(result));
149 }
150}
151
152#[derive(Clone)]
154pub struct WebmcpServer {
155 state: Arc<ServerState>,
156 config: WebmcpServerConfig,
157}
158
159impl std::fmt::Debug for WebmcpServer {
160 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
161 formatter
162 .debug_struct("WebmcpServer")
163 .field("host", &self.config.host)
164 .field("port", &self.config.port)
165 .field("allowed_origins", &self.config.allowed_origins)
166 .finish_non_exhaustive()
167 }
168}
169
170impl WebmcpServer {
171 pub fn new(adapter: Arc<dyn RuntimeAdapter>, config: WebmcpServerConfig) -> Result<Self> {
173 validate_config(&config)?;
174 let pairing = PairingManager::new(&config.allowed_origins, Duration::from_secs(config.pairing_ttl_secs))?;
175 let remote_mcp = config
176 .remote_mcp
177 .clone()
178 .map(|remote_config| RemoteMcpEndpoint::new(Arc::clone(&adapter), remote_config).map(Arc::new))
179 .transpose()?;
180 let event_limit = config
181 .max_frame_bytes
182 .saturating_sub(EVENT_ENVELOPE_OVERHEAD)
183 .min(MAX_EVENT_BYTES);
184 let event_hub = WebmcpEventHub::new_with_max_event_bytes(config.event_hub, event_limit)?;
185 let in_flight = Arc::new(Semaphore::new(config.max_in_flight_requests));
186 let settings = BridgeSettings {
187 host: config.host.clone(),
188 port: config.port,
189 pairing_ttl_secs: config.pairing_ttl_secs,
190 max_frame_bytes: config.max_frame_bytes,
191 max_in_flight_requests: config.max_in_flight_requests,
192 remote_enabled: config.allow_remote,
193 };
194 let dispatch = Arc::new(DispatchState {
195 adapter: Arc::clone(&adapter),
196 pairing: pairing.clone(),
197 event_hub: event_hub.clone(),
198 settings,
199 in_flight: Arc::clone(&in_flight),
200 request_timeout: config.request_timeout,
201 });
202 Ok(Self {
203 state: Arc::new(ServerState {
204 adapter,
205 pairing,
206 event_hub,
207 dispatch,
208 mutation_supervisor: MutationSupervisor::default(),
209 paired_connections: Arc::new(Semaphore::new(MAX_PAIRED_CONNECTIONS)),
210 connections: Arc::new(Semaphore::new(MAX_CONNECTIONS)),
211 max_frame_bytes: config.max_frame_bytes,
212 request_timeout: config.request_timeout,
213 remote_mcp,
214 }),
215 config,
216 })
217 }
218
219 pub fn begin_pairing(&self) -> PairingDisplay {
221 self.state.pairing.begin_pairing()
222 }
223
224 pub fn begin_pairing_for_origin(&self, origin: impl Into<String>) -> Result<PairingDisplay> {
226 self.state.pairing.begin_pairing_for_origin(origin)
227 }
228
229 pub fn replace_pairing_for_origin(&self, origin: impl Into<String>) -> Result<PairingDisplay> {
231 self.state.pairing.replace_pairing_for_origin(origin)
232 }
233
234 pub fn revoke_all_pairings(&self) {
236 self.state.pairing.revoke_all();
237 }
238
239 pub fn event_hub(&self) -> WebmcpEventHub {
241 self.state.event_hub.clone()
242 }
243
244 pub fn router(&self) -> Router {
246 let mut router = Router::new().route("/webmcp", axum::routing::get(websocket_handler));
247 if let Some(remote_mcp) = self.state.remote_mcp.as_ref() {
248 let remote_routes = remote_mcp
249 .routes()
250 .with_state(Arc::clone(remote_mcp))
251 .with_state(self.state.clone());
252 router = router.merge(remote_routes);
253 }
254 router.with_state(self.state.clone())
255 }
256
257 pub async fn bind(&self) -> Result<TcpListener> {
259 TcpListener::bind((self.config.host.as_str(), self.config.port))
260 .await
261 .map_err(WebmcpError::Io)
262 }
263
264 pub async fn serve_listener(&self, listener: TcpListener) -> Result<()> {
266 axum::serve(listener, self.router())
267 .await
268 .map_err(|error| WebmcpError::Adapter(format!("WebMCP listener failed: {error}")))
269 }
270
271 pub async fn serve(&self) -> Result<()> {
273 let listener = self.bind().await?;
274 self.serve_listener(listener).await
275 }
276}
277
278async fn websocket_handler(
279 State(state): State<Arc<ServerState>>,
280 headers: HeaderMap,
281 upgrade: WebSocketUpgrade,
282) -> Response {
283 let Some(origin) = origin_from_headers(&headers) else {
284 return (StatusCode::FORBIDDEN, "WebMCP requires an Origin header").into_response();
285 };
286 if !state.pairing.is_origin_allowed(origin) {
287 return (StatusCode::FORBIDDEN, "WebMCP origin is not allowed").into_response();
288 }
289 let connection_permit = match state.connections.clone().try_acquire_owned() {
290 Ok(permit) => permit,
291 Err(_error) => return (StatusCode::TOO_MANY_REQUESTS, "WebMCP connection limit reached").into_response(),
292 };
293 let origin = origin.to_string();
294 upgrade
295 .max_message_size(state.max_frame_bytes)
296 .max_frame_size(state.max_frame_bytes)
297 .on_upgrade(move |socket| run_socket(socket, state, origin, connection_permit))
298 .into_response()
299}
300
301fn origin_from_headers(headers: &HeaderMap) -> Option<&str> {
302 headers
303 .get(ORIGIN)
304 .and_then(|value| value.to_str().ok())
305 .filter(|origin| !origin.is_empty())
306}
307
308async fn run_socket(
309 mut socket: ws::WebSocket,
310 state: Arc<ServerState>,
311 origin: String,
312 _connection_permit: OwnedSemaphorePermit,
313) {
314 let handshake_deadline = Instant::now() + HANDSHAKE_TIMEOUT;
315 loop {
316 let remaining = handshake_deadline.saturating_duration_since(Instant::now());
317 if remaining.is_zero() {
318 return;
319 }
320 let message = match tokio::time::timeout(remaining, socket.next()).await {
321 Ok(Some(message)) => message,
322 Ok(None) | Err(_) => return,
323 };
324 let Ok(message) = message else { return };
325 match message {
326 ws::Message::Text(text) => {
327 if text.len() > state.max_frame_bytes {
328 drop(
329 send_response(
330 &mut socket,
331 BridgeResponse::failure("unknown", "frame_too_large", "request exceeds the frame limit"),
332 state.max_frame_bytes,
333 )
334 .await,
335 );
336 return;
337 }
338 match serde_json::from_slice::<BridgeRequest>(text.as_bytes()) {
339 Ok(BridgeRequest::Pair {
340 request_id,
341 code,
342 resume_token,
343 origin: claimed_origin,
344 after_sequence,
345 }) => {
346 if !is_valid_request_id(&request_id) {
347 drop(
348 send_response(
349 &mut socket,
350 invalid_request_id_response(&request_id),
351 state.max_frame_bytes,
352 )
353 .await,
354 );
355 continue;
356 }
357 if claimed_origin.as_deref().is_some_and(|claimed| claimed != origin) {
358 drop(
359 send_response(
360 &mut socket,
361 BridgeResponse::failure(
362 &request_id,
363 "origin_mismatch",
364 "request origin does not match the WebSocket origin",
365 ),
366 state.max_frame_bytes,
367 )
368 .await,
369 );
370 return;
371 }
372 let mut subscription = match state.event_hub.subscribe(after_sequence) {
373 Ok(subscription) => subscription,
374 Err(error) => {
375 drop(
376 send_response(
377 &mut socket,
378 response_for_error(&request_id, error),
379 state.max_frame_bytes,
380 )
381 .await,
382 );
383 return;
384 }
385 };
386 let connection_permit = match state.paired_connections.clone().try_acquire_owned() {
387 Ok(permit) => permit,
388 Err(_error) => {
389 drop(
390 send_response(
391 &mut socket,
392 BridgeResponse::failure(
393 &request_id,
394 "connection_limit",
395 "the WebMCP server has reached its paired connection limit",
396 ),
397 state.max_frame_bytes,
398 )
399 .await,
400 );
401 return;
402 }
403 };
404 let session = match resume_token {
405 Some(token) => state.pairing.resume(&token, &origin),
406 None => state.pairing.pair(&code, &origin),
407 };
408 match session {
409 Ok(session) => {
410 let response = BridgeResponse::success(
411 request_id,
412 PairPayload {
413 token: session.token().to_string(),
414 protocol_version: PROTOCOL_VERSION,
415 expires_in_secs: session.expires_in().as_secs().max(1),
416 },
417 );
418 if send_response(&mut socket, response, state.max_frame_bytes).await.is_err() {
419 return;
420 }
421 for event in subscription.replay() {
422 if send_event(&mut socket, event.sequence, &event.event, state.max_frame_bytes)
423 .await
424 .is_err()
425 {
426 return;
427 }
428 }
429 run_paired_socket(
430 socket,
431 state,
432 origin,
433 session.token().to_string(),
434 &mut subscription,
435 connection_permit,
436 _connection_permit,
437 )
438 .await;
439 return;
440 }
441 Err(error) => {
442 drop(
443 send_response(
444 &mut socket,
445 response_for_error(&request_id, error),
446 state.max_frame_bytes,
447 )
448 .await,
449 );
450 }
451 }
452 }
453 Ok(request) => {
454 let request_id = request.request_id().to_string();
455 let response = if is_valid_request_id(&request_id) {
456 response_for_error(&request_id, WebmcpError::Unauthorized)
457 } else {
458 invalid_request_id_response(&request_id)
459 };
460 drop(send_response(&mut socket, response, state.max_frame_bytes).await);
461 }
462 Err(error) => {
463 drop(
464 send_response(
465 &mut socket,
466 BridgeResponse::failure("unknown", "malformed_request", error.to_string()),
467 state.max_frame_bytes,
468 )
469 .await,
470 );
471 }
472 }
473 }
474 ws::Message::Binary(_) => {
475 drop(
476 send_response(
477 &mut socket,
478 BridgeResponse::failure(
479 "unknown",
480 "binary_not_supported",
481 "WebMCP accepts JSON text frames only",
482 ),
483 state.max_frame_bytes,
484 )
485 .await,
486 );
487 }
488 ws::Message::Ping(payload) => {
489 if send_pong(&mut socket, payload).await.is_err() {
490 return;
491 }
492 }
493 ws::Message::Close(_) => return,
494 ws::Message::Pong(_) => {}
495 }
496 }
497}
498
499enum PairedFrameAction {
500 Request(BridgeRequest),
501 Continue,
502 Close,
503}
504
505async fn handle_paired_frame(
506 socket: &mut ws::WebSocket,
507 message: ws::Message,
508 max_frame_bytes: usize,
509) -> PairedFrameAction {
510 match message {
511 ws::Message::Text(text) => {
512 if text.len() > max_frame_bytes {
513 drop(
514 send_response(
515 socket,
516 BridgeResponse::failure("unknown", "frame_too_large", "request exceeds the frame limit"),
517 max_frame_bytes,
518 )
519 .await,
520 );
521 return PairedFrameAction::Close;
522 }
523 let request = match serde_json::from_slice::<BridgeRequest>(text.as_bytes()) {
524 Ok(request) => request,
525 Err(error) => {
526 if send_response(
527 socket,
528 BridgeResponse::failure("unknown", "malformed_request", error.to_string()),
529 max_frame_bytes,
530 )
531 .await
532 .is_err()
533 {
534 return PairedFrameAction::Close;
535 }
536 return PairedFrameAction::Continue;
537 }
538 };
539 if !is_valid_request_id(request.request_id()) {
540 if send_response(socket, invalid_request_id_response(request.request_id()), max_frame_bytes)
541 .await
542 .is_err()
543 {
544 return PairedFrameAction::Close;
545 }
546 return PairedFrameAction::Continue;
547 }
548 PairedFrameAction::Request(request)
549 }
550 ws::Message::Binary(_) => {
551 if send_response(
552 socket,
553 BridgeResponse::failure("unknown", "binary_not_supported", "WebMCP accepts JSON text frames only"),
554 max_frame_bytes,
555 )
556 .await
557 .is_err()
558 {
559 PairedFrameAction::Close
560 } else {
561 PairedFrameAction::Continue
562 }
563 }
564 ws::Message::Ping(payload) => {
565 if send_pong(socket, payload).await.is_err() {
566 PairedFrameAction::Close
567 } else {
568 PairedFrameAction::Continue
569 }
570 }
571 ws::Message::Close(_) => PairedFrameAction::Close,
572 ws::Message::Pong(_) => PairedFrameAction::Continue,
573 }
574}
575
576async fn run_paired_socket(
577 mut socket: ws::WebSocket,
578 state: Arc<ServerState>,
579 origin: String,
580 token: String,
581 subscription: &mut EventHubSubscription,
582 _connection_permit: OwnedSemaphorePermit,
583 _all_connections_permit: OwnedSemaphorePermit,
584) {
585 let mut expiry_check = tokio::time::interval(SESSION_LEASE_CHECK_INTERVAL);
586 expiry_check.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
587 loop {
588 tokio::select! {
589 message = socket.next() => {
590 let Some(Ok(message)) = message else { return };
591 match handle_paired_frame(&mut socket, message, state.max_frame_bytes).await {
592 PairedFrameAction::Close => return,
593 PairedFrameAction::Continue => {}
594 PairedFrameAction::Request(request) => {
595 let request_id = request.request_id().to_string();
596 let mutation_request = matches!(
597 &request,
598 BridgeRequest::ApplyProposal { .. } | BridgeRequest::RevertLastChange { .. }
599 );
600 let operation: BoxFuture<'static, Result<Value>> = if mutation_request {
601 let supervisor = state.mutation_supervisor.clone();
602 let dispatch = Arc::clone(&state.dispatch);
603 let origin = origin.clone();
604 let token = token.clone();
605 Box::pin(async move { supervisor.submit(dispatch, origin, token, request).await })
606 } else {
607 Box::pin(dispatch_request(
608 Arc::clone(&state.dispatch),
609 origin.clone(),
610 token.clone(),
611 request,
612 ))
613 };
614 let mut operation = operation;
615 loop {
616 tokio::select! {
617 result = &mut operation => {
618 let response = match result {
619 Ok(payload) => BridgeResponse::success(request_id.clone(), payload),
620 Err(error) => response_for_error(&request_id, error),
621 };
622 if send_response(&mut socket, response, state.max_frame_bytes).await.is_err() { return; }
623 break;
624 }
625 message = socket.next() => {
626 let Some(Ok(message)) = message else { return };
627 match handle_paired_frame(&mut socket, message, state.max_frame_bytes).await {
628 PairedFrameAction::Close => return,
629 PairedFrameAction::Continue => {}
630 PairedFrameAction::Request(request) => {
631 let request_id = request.request_id().to_string();
632 let response = if request.token() != Some(token.as_str()) {
633 response_for_error(&request_id, WebmcpError::Unauthorized)
634 } else if matches!(&request, BridgeRequest::Cancel { .. }) {
635 match dispatch_cancel_request(&state, &origin, &token, request).await {
636 Ok(payload) => BridgeResponse::success(request_id.clone(), payload),
637 Err(error) => response_for_error(&request_id, error),
638 }
639 } else {
640 if state.pairing.refresh(&token, &origin).is_err() { return; }
641 BridgeResponse::failure(&request_id, "request_in_progress", "wait for the active request or cancel it")
642 };
643 if send_response(&mut socket, response, state.max_frame_bytes).await.is_err() { return; }
644 }
645 }
646 }
647 event = subscription.recv() => {
648 let Some(event) = event else {
649 drop(send_response(&mut socket, BridgeResponse::failure("event", "slow_client", "client could not keep up with runtime events"), state.max_frame_bytes).await);
650 return;
651 };
652 if state.pairing.validate(&token, &origin).is_err() { return; }
653 if send_event(&mut socket, event.sequence, &event.event, state.max_frame_bytes).await.is_err() { return; }
654 }
655 _ = expiry_check.tick() => {
656 if state.pairing.refresh(&token, &origin).is_err() { return; }
660 }
661 }
662 }
663 }
664 }
665 }
666 event = subscription.recv() => {
667 let Some(event) = event else {
668 drop(send_response(&mut socket, BridgeResponse::failure("event", "slow_client", "client could not keep up with runtime events"), state.max_frame_bytes).await);
669 return;
670 };
671 if state.pairing.validate(&token, &origin).is_err() { return; }
672 if send_event(&mut socket, event.sequence, &event.event, state.max_frame_bytes).await.is_err() { return; }
673 }
674 _ = expiry_check.tick() => {
675 if state.pairing.validate(&token, &origin).is_err() { return; }
676 }
677 }
678 }
679}
680
681async fn dispatch_request(
682 dispatch: Arc<DispatchState>,
683 origin: String,
684 token: String,
685 request: BridgeRequest,
686) -> Result<Value> {
687 if request.token() != Some(token.as_str()) {
688 return Err(WebmcpError::Unauthorized);
689 }
690 dispatch.pairing.refresh(&token, &origin)?;
691 let _permit = tokio::time::timeout(dispatch.request_timeout, dispatch.in_flight.clone().acquire_owned())
692 .await
693 .map_err(|_error| WebmcpError::Timeout(dispatch.request_timeout))?
694 .map_err(|_error| WebmcpError::Adapter("WebMCP request capacity is closed".to_string()))?;
695 dispatch.pairing.refresh(&token, &origin)?;
699 let mutation_request =
700 matches!(&request, BridgeRequest::ApplyProposal { .. } | BridgeRequest::RevertLastChange { .. });
701 let operation = async {
702 match request {
703 BridgeRequest::Pair { .. } => Err(WebmcpError::Unauthorized),
704 BridgeRequest::ExplanationGet { scope, offset, .. } => {
705 serde_json::to_value(dispatch.adapter.explanation_get(scope, offset).await?).map_err(WebmcpError::from)
706 }
707 BridgeRequest::ExplanationEvidence { reference, offset, .. } => {
708 serde_json::to_value(dispatch.adapter.explanation_evidence(reference, offset).await?)
709 .map_err(WebmcpError::from)
710 }
711 BridgeRequest::ExplanationNavigate { reference, .. } => {
712 Ok(serde_json::json!({"focused": dispatch.adapter.explanation_navigate(reference).await?}))
713 }
714 BridgeRequest::Status { .. } => {
715 let runtime = dispatch.adapter.status().await?;
716 serde_json::to_value(StatusPayload {
717 protocol_version: PROTOCOL_VERSION,
718 connected: runtime.connected,
719 runtime,
720 authenticated_origin: origin,
721 settings: dispatch.settings.clone(),
722 latest_sequence: dispatch.event_hub.latest_sequence(),
723 })
724 .map_err(WebmcpError::Json)
725 }
726 BridgeRequest::ListFiles { .. } => {
727 serde_json::to_value(dispatch.adapter.list_files().await?).map_err(WebmcpError::Json)
728 }
729 BridgeRequest::ReadFile { path, .. } => {
730 serde_json::to_value(dispatch.adapter.read_file(&path).await?).map_err(WebmcpError::Json)
731 }
732 BridgeRequest::ProposeChanges { changes, .. } => {
733 serde_json::to_value(dispatch.adapter.propose_changes(changes).await?).map_err(WebmcpError::Json)
734 }
735 BridgeRequest::ApplyProposal { proposal_id, .. } => {
736 serde_json::to_value(dispatch.adapter.apply_proposal(&proposal_id).await?).map_err(WebmcpError::Json)
737 }
738 BridgeRequest::RunChecks { command, .. } => {
739 serde_json::to_value(dispatch.adapter.run_checks(&command).await?).map_err(WebmcpError::Json)
740 }
741 BridgeRequest::RevertLastChange { change_id, .. } => {
742 serde_json::to_value(dispatch.adapter.revert_last_change(&change_id).await?).map_err(WebmcpError::Json)
743 }
744 BridgeRequest::RequestTurn { prompt, proposal_id, .. } => {
745 serde_json::to_value(dispatch.adapter.request_turn(&prompt, proposal_id.as_deref()).await?)
746 .map_err(WebmcpError::Json)
747 }
748 BridgeRequest::Cancel { target_id, .. } => {
749 let accepted = dispatch.adapter.cancel(&target_id).await?;
750 Ok(serde_json::json!({ "cancelled": target_id, "accepted": accepted }))
751 }
752 }
753 };
754 if mutation_request {
755 operation.await
759 } else {
760 tokio::time::timeout(dispatch.request_timeout, operation)
761 .await
762 .map_err(|_error| WebmcpError::Timeout(dispatch.request_timeout))?
763 }
764}
765
766async fn dispatch_cancel_request(
767 state: &Arc<ServerState>,
768 origin: &str,
769 token: &str,
770 request: BridgeRequest,
771) -> Result<Value> {
772 if request.token() != Some(token) {
773 return Err(WebmcpError::Unauthorized);
774 }
775 state.pairing.refresh(token, origin)?;
776 let BridgeRequest::Cancel { target_id, .. } = request else {
777 return Err(WebmcpError::InvalidRequest(
778 "only cancellation requests are accepted while a request is running".to_string(),
779 ));
780 };
781 let accepted = tokio::time::timeout(state.request_timeout, state.adapter.cancel(&target_id))
782 .await
783 .map_err(|_error| WebmcpError::Timeout(state.request_timeout))??;
784 Ok(serde_json::json!({ "cancelled": target_id, "accepted": accepted }))
785}
786
787async fn send_response(socket: &mut ws::WebSocket, response: BridgeResponse, max_frame_bytes: usize) -> Result<()> {
788 let serialized = serde_json::to_string(&response)?;
789 let serialized = if serialized.len() > max_frame_bytes {
790 serde_json::to_string(&BridgeResponse::failure(
791 response.request_id,
792 "limit_exceeded",
793 "WebMCP response exceeds the configured frame limit",
794 ))?
795 } else {
796 serialized
797 };
798 if serialized.len() > max_frame_bytes {
799 return Err(WebmcpError::LimitExceeded);
800 }
801 tokio::time::timeout(WRITE_TIMEOUT, socket.send(ws::Message::Text(serialized.into())))
802 .await
803 .map_err(|_error| WebmcpError::Timeout(WRITE_TIMEOUT))?
804 .map_err(|error| WebmcpError::Adapter(error.to_string()))
805}
806
807async fn send_pong(socket: &mut ws::WebSocket, payload: axum::body::Bytes) -> Result<()> {
808 tokio::time::timeout(WRITE_TIMEOUT, socket.send(ws::Message::Pong(payload)))
809 .await
810 .map_err(|_error| WebmcpError::Timeout(WRITE_TIMEOUT))?
811 .map_err(|error| WebmcpError::Adapter(error.to_string()))
812}
813
814async fn send_event(
815 socket: &mut ws::WebSocket,
816 sequence: u64,
817 event: &vtcode_exec_events::VersionedThreadEvent,
818 max_frame_bytes: usize,
819) -> Result<()> {
820 let message = BridgeEventMessage { kind: "event", sequence, event: event.clone() };
821 let serialized = serde_json::to_string(&message)?;
822 if serialized.len() > max_frame_bytes {
823 return Err(WebmcpError::LimitExceeded);
824 }
825 tokio::time::timeout(WRITE_TIMEOUT, socket.send(ws::Message::Text(serialized.into())))
826 .await
827 .map_err(|_error| WebmcpError::Timeout(WRITE_TIMEOUT))?
828 .map_err(|error| WebmcpError::Adapter(error.to_string()))
829}
830
831fn response_for_error(request_id: &str, error: WebmcpError) -> BridgeResponse {
832 let (code, message) = match &error {
833 WebmcpError::OriginRejected(_) => ("origin_rejected", "browser origin is not allowed".to_string()),
834 WebmcpError::PairingExpired => ("pairing_expired", error.to_string()),
835 WebmcpError::PairingUsed => ("pairing_used", error.to_string()),
836 WebmcpError::Unauthorized => ("unauthorized", error.to_string()),
837 WebmcpError::LimitExceeded => ("limit_exceeded", error.to_string()),
838 WebmcpError::PathRejected(_) => ("path_rejected", error.to_string()),
839 WebmcpError::Conflict { .. } => ("conflict", error.to_string()),
840 WebmcpError::ProposalNotFound => ("proposal_not_found", error.to_string()),
841 WebmcpError::ApprovalRequired => ("approval_required", error.to_string()),
842 WebmcpError::Unsupported(_) => ("unsupported", error.to_string()),
843 WebmcpError::ChangeNotFound => ("change_not_found", error.to_string()),
844 WebmcpError::PartialApply => ("partial_apply", error.to_string()),
845 WebmcpError::SequenceGap { .. } => ("sequence_gap", error.to_string()),
846 WebmcpError::SlowClient => ("slow_client", error.to_string()),
847 WebmcpError::Timeout(_) => ("timeout", error.to_string()),
848 WebmcpError::InvalidRequest(_) | WebmcpError::Json(_) => ("invalid_request", error.to_string()),
849 WebmcpError::Io(_) | WebmcpError::Adapter(_) => {
850 ("runtime_error", "WebMCP runtime operation failed".to_string())
851 }
852 };
853 BridgeResponse::failure(response_request_id(request_id), code, message)
854}
855
856fn invalid_request_id_response(request_id: &str) -> BridgeResponse {
857 BridgeResponse::failure(
858 response_request_id(request_id),
859 "invalid_request",
860 "request_id must be between 1 and 256 UTF-8 bytes",
861 )
862}
863
864fn validate_config(config: &WebmcpServerConfig) -> Result<()> {
865 if config.host.trim().is_empty() || config.max_in_flight_requests == 0 {
866 return Err(WebmcpError::InvalidRequest("WebMCP host and limits must be non-empty".to_string()));
867 }
868 if config.max_frame_bytes < MIN_FRAME_BYTES
869 || config.max_frame_bytes > 16 * 1024 * 1024
870 || config.max_in_flight_requests > 64
871 {
872 return Err(WebmcpError::LimitExceeded);
873 }
874 if (config.allowed_origins.is_empty() && config.remote_mcp.is_none())
875 || config.allowed_origins.iter().any(|origin| !is_valid_origin(origin))
876 {
877 return Err(WebmcpError::InvalidRequest("WebMCP requires an explicit origin allowlist".to_string()));
878 }
879 let address = config
880 .host
881 .parse::<IpAddr>()
882 .map_err(|_error| WebmcpError::InvalidRequest("WebMCP host must be a literal IP address".to_string()))?;
883 if !address.is_loopback() {
884 return Err(WebmcpError::InvalidRequest(
885 "WebMCP only binds loopback; place a TLS-terminating reverse proxy in front of it for remote access"
886 .to_string(),
887 ));
888 }
889 match (config.allow_remote, config.public_url.as_deref()) {
890 (false, Some(_)) => {
891 return Err(WebmcpError::InvalidRequest("--public-url requires remote WebMCP mode".to_string()));
892 }
893 (true, None) => {
894 return Err(WebmcpError::InvalidRequest("remote WebMCP mode requires a wss:// public URL".to_string()));
895 }
896 (true, Some(url)) if !is_valid_public_url(url) => {
897 return Err(WebmcpError::InvalidRequest(
898 "remote WebMCP mode requires a valid wss:// public URL".to_string(),
899 ));
900 }
901 (false, None) => {}
902 (true, Some(_)) => {}
903 }
904 if config.request_timeout.is_zero() {
905 return Err(WebmcpError::InvalidRequest("WebMCP request timeout must be positive".to_string()));
906 }
907 if let Some(remote_mcp) = config.remote_mcp.as_ref() {
908 remote_mcp.validate()?;
909 }
910 Ok(())
911}
912
913fn is_valid_public_url(url: &str) -> bool {
914 let Ok(parsed) = url::Url::parse(url) else {
915 return false;
916 };
917 url == url.trim()
918 && !url.chars().any(char::is_whitespace)
919 && parsed.scheme() == "wss"
920 && parsed.host_str().is_some_and(|host| !host.is_empty())
921 && parsed.username().is_empty()
922 && parsed.password().is_none()
923 && parsed.query().is_none()
924 && parsed.fragment().is_none()
925}
926
927#[cfg(test)]
928mod tests {
929 use super::*;
930 use crate::FilesystemWorkspace;
931 use tempfile::TempDir;
932
933 #[tokio::test]
934 async fn explanation_requests_require_origin_bound_tokens_and_report_unsupported() {
935 let temp = TempDir::new().expect("workspace");
936 let adapter = Arc::new(FilesystemWorkspace::new(temp.path(), [], false).await.expect("adapter"));
937 let origin = "https://example.test";
938 let server = WebmcpServer::new(
939 adapter,
940 WebmcpServerConfig {
941 allowed_origins: vec![origin.to_owned(), "https://other.test".to_owned()],
942 ..Default::default()
943 },
944 )
945 .expect("server");
946 let pairing = server.begin_pairing_for_origin(origin).expect("pairing");
947 let session = server
948 .state
949 .pairing
950 .pair(pairing.code(), origin)
951 .expect("authenticated session");
952 let token = session.token().to_owned();
953 let reference = vtcode_memory::explanation::EvidenceRef {
954 session_id: "session".into(),
955 offset: 0,
956 length: 1,
957 digest: "0".repeat(64),
958 item_id: None,
959 };
960 let requests = [
961 BridgeRequest::ExplanationGet {
962 request_id: "get".into(),
963 token: token.clone(),
964 scope: vtcode_memory::explanation::ExplanationScope::Task,
965 offset: 0,
966 },
967 BridgeRequest::ExplanationEvidence {
968 request_id: "evidence".into(),
969 token: token.clone(),
970 reference: reference.clone(),
971 offset: 0,
972 },
973 BridgeRequest::ExplanationNavigate {
974 request_id: "navigate".into(),
975 token: token.clone(),
976 reference,
977 },
978 ];
979 for request in &requests {
980 let dispatch = || Arc::clone(&server.state.dispatch);
981 assert!(matches!(
982 dispatch_request(dispatch(), origin.into(), "wrong-token".into(), request.clone()).await,
983 Err(WebmcpError::Unauthorized)
984 ));
985 let error = dispatch_request(dispatch(), origin.into(), token.clone(), request.clone())
986 .await
987 .expect_err("headless adapter does not provide explanations");
988 assert!(matches!(error, WebmcpError::Unsupported(_)));
989 assert_eq!(response_for_error("operation", error).error.expect("error payload").code, "unsupported");
990 }
991 assert!(matches!(
992 dispatch_request(
993 Arc::clone(&server.state.dispatch),
994 "https://other.test".into(),
995 token.clone(),
996 requests[0].clone()
997 )
998 .await,
999 Err(WebmcpError::Unauthorized)
1000 ));
1001 for request in requests {
1002 assert!(matches!(
1003 dispatch_request(Arc::clone(&server.state.dispatch), origin.into(), token.clone(), request).await,
1004 Err(WebmcpError::Unauthorized)
1005 ));
1006 }
1007 let status = server.state.adapter.status().await.expect("status");
1008 assert!(!status.explanations_available);
1009 assert!(!status.turns_available);
1010 assert!(std::fs::read_dir(temp.path()).expect("workspace files").next().is_none());
1011 }
1012
1013 #[tokio::test]
1014 async fn server_requires_explicit_origins_and_remote_flags() {
1015 let temp = TempDir::new().expect("temp dir");
1016 let adapter = Arc::new(FilesystemWorkspace::new(temp.path(), [], false).await.expect("adapter"));
1017 assert!(matches!(
1018 WebmcpServer::new(adapter.clone(), WebmcpServerConfig::default()),
1019 Err(WebmcpError::InvalidRequest(_))
1020 ));
1021 let config = WebmcpServerConfig {
1022 host: "0.0.0.0".to_string(),
1023 allowed_origins: vec!["https://example.test".to_string()],
1024 ..Default::default()
1025 };
1026 assert!(matches!(WebmcpServer::new(adapter, config), Err(WebmcpError::InvalidRequest(_))));
1027
1028 let remote_config = WebmcpServerConfig {
1029 host: "0.0.0.0".to_string(),
1030 allowed_origins: vec!["https://example.test".to_string()],
1031 allow_remote: true,
1032 public_url: Some("wss://bridge.example.test/webmcp".to_string()),
1033 ..Default::default()
1034 };
1035 let temp = TempDir::new().expect("temp dir");
1036 let adapter = Arc::new(FilesystemWorkspace::new(temp.path(), [], false).await.expect("adapter"));
1037 assert!(matches!(WebmcpServer::new(adapter, remote_config), Err(WebmcpError::InvalidRequest(_))));
1038
1039 let temp = TempDir::new().expect("temp dir");
1040 let adapter = Arc::new(FilesystemWorkspace::new(temp.path(), [], false).await.expect("adapter"));
1041 let valid_proxy_config = WebmcpServerConfig {
1042 allowed_origins: vec!["https://example.test".to_string()],
1043 allow_remote: true,
1044 public_url: Some("wss://bridge.example.test/webmcp".to_string()),
1045 ..Default::default()
1046 };
1047 assert!(WebmcpServer::new(adapter.clone(), valid_proxy_config).is_ok());
1048 let invalid_public_url_config = WebmcpServerConfig {
1049 allowed_origins: vec!["https://example.test".to_string()],
1050 public_url: Some("ws://bridge.example.test/webmcp".to_string()),
1051 ..Default::default()
1052 };
1053 assert!(matches!(WebmcpServer::new(adapter, invalid_public_url_config), Err(WebmcpError::InvalidRequest(_))));
1054
1055 let temp = TempDir::new().expect("temp dir");
1056 let adapter = Arc::new(FilesystemWorkspace::new(temp.path(), [], false).await.expect("adapter"));
1057 let invalid_public_url_config = WebmcpServerConfig {
1058 allowed_origins: vec!["https://example.test".to_string()],
1059 allow_remote: true,
1060 public_url: Some("wss://".to_string()),
1061 ..Default::default()
1062 };
1063 assert!(matches!(WebmcpServer::new(adapter, invalid_public_url_config), Err(WebmcpError::InvalidRequest(_))));
1064 }
1065}