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::Status { .. } => {
705 let runtime = dispatch.adapter.status().await?;
706 serde_json::to_value(StatusPayload {
707 protocol_version: PROTOCOL_VERSION,
708 connected: runtime.connected,
709 runtime,
710 authenticated_origin: origin,
711 settings: dispatch.settings.clone(),
712 latest_sequence: dispatch.event_hub.latest_sequence(),
713 })
714 .map_err(WebmcpError::Json)
715 }
716 BridgeRequest::ListFiles { .. } => {
717 serde_json::to_value(dispatch.adapter.list_files().await?).map_err(WebmcpError::Json)
718 }
719 BridgeRequest::ReadFile { path, .. } => {
720 serde_json::to_value(dispatch.adapter.read_file(&path).await?).map_err(WebmcpError::Json)
721 }
722 BridgeRequest::ProposeChanges { changes, .. } => {
723 serde_json::to_value(dispatch.adapter.propose_changes(changes).await?).map_err(WebmcpError::Json)
724 }
725 BridgeRequest::ApplyProposal { proposal_id, .. } => {
726 serde_json::to_value(dispatch.adapter.apply_proposal(&proposal_id).await?).map_err(WebmcpError::Json)
727 }
728 BridgeRequest::RunChecks { command, .. } => {
729 serde_json::to_value(dispatch.adapter.run_checks(&command).await?).map_err(WebmcpError::Json)
730 }
731 BridgeRequest::RevertLastChange { change_id, .. } => {
732 serde_json::to_value(dispatch.adapter.revert_last_change(&change_id).await?).map_err(WebmcpError::Json)
733 }
734 BridgeRequest::RequestTurn { prompt, proposal_id, .. } => {
735 serde_json::to_value(dispatch.adapter.request_turn(&prompt, proposal_id.as_deref()).await?)
736 .map_err(WebmcpError::Json)
737 }
738 BridgeRequest::Cancel { target_id, .. } => {
739 let accepted = dispatch.adapter.cancel(&target_id).await?;
740 Ok(serde_json::json!({ "cancelled": target_id, "accepted": accepted }))
741 }
742 }
743 };
744 if mutation_request {
745 operation.await
749 } else {
750 tokio::time::timeout(dispatch.request_timeout, operation)
751 .await
752 .map_err(|_error| WebmcpError::Timeout(dispatch.request_timeout))?
753 }
754}
755
756async fn dispatch_cancel_request(
757 state: &Arc<ServerState>,
758 origin: &str,
759 token: &str,
760 request: BridgeRequest,
761) -> Result<Value> {
762 if request.token() != Some(token) {
763 return Err(WebmcpError::Unauthorized);
764 }
765 state.pairing.refresh(token, origin)?;
766 let BridgeRequest::Cancel { target_id, .. } = request else {
767 return Err(WebmcpError::InvalidRequest(
768 "only cancellation requests are accepted while a request is running".to_string(),
769 ));
770 };
771 let accepted = tokio::time::timeout(state.request_timeout, state.adapter.cancel(&target_id))
772 .await
773 .map_err(|_error| WebmcpError::Timeout(state.request_timeout))??;
774 Ok(serde_json::json!({ "cancelled": target_id, "accepted": accepted }))
775}
776
777async fn send_response(socket: &mut ws::WebSocket, response: BridgeResponse, max_frame_bytes: usize) -> Result<()> {
778 let serialized = serde_json::to_string(&response)?;
779 let serialized = if serialized.len() > max_frame_bytes {
780 serde_json::to_string(&BridgeResponse::failure(
781 response.request_id,
782 "limit_exceeded",
783 "WebMCP response exceeds the configured frame limit",
784 ))?
785 } else {
786 serialized
787 };
788 if serialized.len() > max_frame_bytes {
789 return Err(WebmcpError::LimitExceeded);
790 }
791 tokio::time::timeout(WRITE_TIMEOUT, socket.send(ws::Message::Text(serialized.into())))
792 .await
793 .map_err(|_error| WebmcpError::Timeout(WRITE_TIMEOUT))?
794 .map_err(|error| WebmcpError::Adapter(error.to_string()))
795}
796
797async fn send_pong(socket: &mut ws::WebSocket, payload: axum::body::Bytes) -> Result<()> {
798 tokio::time::timeout(WRITE_TIMEOUT, socket.send(ws::Message::Pong(payload)))
799 .await
800 .map_err(|_error| WebmcpError::Timeout(WRITE_TIMEOUT))?
801 .map_err(|error| WebmcpError::Adapter(error.to_string()))
802}
803
804async fn send_event(
805 socket: &mut ws::WebSocket,
806 sequence: u64,
807 event: &vtcode_exec_events::VersionedThreadEvent,
808 max_frame_bytes: usize,
809) -> Result<()> {
810 let message = BridgeEventMessage { kind: "event", sequence, event: event.clone() };
811 let serialized = serde_json::to_string(&message)?;
812 if serialized.len() > max_frame_bytes {
813 return Err(WebmcpError::LimitExceeded);
814 }
815 tokio::time::timeout(WRITE_TIMEOUT, socket.send(ws::Message::Text(serialized.into())))
816 .await
817 .map_err(|_error| WebmcpError::Timeout(WRITE_TIMEOUT))?
818 .map_err(|error| WebmcpError::Adapter(error.to_string()))
819}
820
821fn response_for_error(request_id: &str, error: WebmcpError) -> BridgeResponse {
822 let (code, message) = match &error {
823 WebmcpError::OriginRejected(_) => ("origin_rejected", "browser origin is not allowed".to_string()),
824 WebmcpError::PairingExpired => ("pairing_expired", error.to_string()),
825 WebmcpError::PairingUsed => ("pairing_used", error.to_string()),
826 WebmcpError::Unauthorized => ("unauthorized", error.to_string()),
827 WebmcpError::LimitExceeded => ("limit_exceeded", error.to_string()),
828 WebmcpError::PathRejected(_) => ("path_rejected", error.to_string()),
829 WebmcpError::Conflict { .. } => ("conflict", error.to_string()),
830 WebmcpError::ProposalNotFound => ("proposal_not_found", error.to_string()),
831 WebmcpError::ApprovalRequired => ("approval_required", error.to_string()),
832 WebmcpError::Unsupported(_) => ("unsupported", error.to_string()),
833 WebmcpError::ChangeNotFound => ("change_not_found", error.to_string()),
834 WebmcpError::PartialApply => ("partial_apply", error.to_string()),
835 WebmcpError::SequenceGap { .. } => ("sequence_gap", error.to_string()),
836 WebmcpError::SlowClient => ("slow_client", error.to_string()),
837 WebmcpError::Timeout(_) => ("timeout", error.to_string()),
838 WebmcpError::InvalidRequest(_) | WebmcpError::Json(_) => ("invalid_request", error.to_string()),
839 WebmcpError::Io(_) | WebmcpError::Adapter(_) => {
840 ("runtime_error", "WebMCP runtime operation failed".to_string())
841 }
842 };
843 BridgeResponse::failure(response_request_id(request_id), code, message)
844}
845
846fn invalid_request_id_response(request_id: &str) -> BridgeResponse {
847 BridgeResponse::failure(
848 response_request_id(request_id),
849 "invalid_request",
850 "request_id must be between 1 and 256 UTF-8 bytes",
851 )
852}
853
854fn validate_config(config: &WebmcpServerConfig) -> Result<()> {
855 if config.host.trim().is_empty() || config.max_in_flight_requests == 0 {
856 return Err(WebmcpError::InvalidRequest("WebMCP host and limits must be non-empty".to_string()));
857 }
858 if config.max_frame_bytes < MIN_FRAME_BYTES
859 || config.max_frame_bytes > 16 * 1024 * 1024
860 || config.max_in_flight_requests > 64
861 {
862 return Err(WebmcpError::LimitExceeded);
863 }
864 if (config.allowed_origins.is_empty() && config.remote_mcp.is_none())
865 || config.allowed_origins.iter().any(|origin| !is_valid_origin(origin))
866 {
867 return Err(WebmcpError::InvalidRequest("WebMCP requires an explicit origin allowlist".to_string()));
868 }
869 let address = config
870 .host
871 .parse::<IpAddr>()
872 .map_err(|_error| WebmcpError::InvalidRequest("WebMCP host must be a literal IP address".to_string()))?;
873 if !address.is_loopback() {
874 return Err(WebmcpError::InvalidRequest(
875 "WebMCP only binds loopback; place a TLS-terminating reverse proxy in front of it for remote access"
876 .to_string(),
877 ));
878 }
879 match (config.allow_remote, config.public_url.as_deref()) {
880 (false, Some(_)) => {
881 return Err(WebmcpError::InvalidRequest("--public-url requires remote WebMCP mode".to_string()));
882 }
883 (true, None) => {
884 return Err(WebmcpError::InvalidRequest("remote WebMCP mode requires a wss:// public URL".to_string()));
885 }
886 (true, Some(url)) if !is_valid_public_url(url) => {
887 return Err(WebmcpError::InvalidRequest(
888 "remote WebMCP mode requires a valid wss:// public URL".to_string(),
889 ));
890 }
891 (false, None) => {}
892 (true, Some(_)) => {}
893 }
894 if config.request_timeout.is_zero() {
895 return Err(WebmcpError::InvalidRequest("WebMCP request timeout must be positive".to_string()));
896 }
897 if let Some(remote_mcp) = config.remote_mcp.as_ref() {
898 remote_mcp.validate()?;
899 }
900 Ok(())
901}
902
903fn is_valid_public_url(url: &str) -> bool {
904 let Ok(parsed) = url::Url::parse(url) else {
905 return false;
906 };
907 url == url.trim()
908 && !url.chars().any(char::is_whitespace)
909 && parsed.scheme() == "wss"
910 && parsed.host_str().is_some_and(|host| !host.is_empty())
911 && parsed.username().is_empty()
912 && parsed.password().is_none()
913 && parsed.query().is_none()
914 && parsed.fragment().is_none()
915}
916
917#[cfg(test)]
918mod tests {
919 use super::*;
920 use crate::FilesystemWorkspace;
921 use tempfile::TempDir;
922
923 #[tokio::test]
924 async fn server_requires_explicit_origins_and_remote_flags() {
925 let temp = TempDir::new().expect("temp dir");
926 let adapter = Arc::new(FilesystemWorkspace::new(temp.path(), [], false).await.expect("adapter"));
927 assert!(matches!(
928 WebmcpServer::new(adapter.clone(), WebmcpServerConfig::default()),
929 Err(WebmcpError::InvalidRequest(_))
930 ));
931 let config = WebmcpServerConfig {
932 host: "0.0.0.0".to_string(),
933 allowed_origins: vec!["https://example.test".to_string()],
934 ..Default::default()
935 };
936 assert!(matches!(WebmcpServer::new(adapter, config), Err(WebmcpError::InvalidRequest(_))));
937
938 let remote_config = WebmcpServerConfig {
939 host: "0.0.0.0".to_string(),
940 allowed_origins: vec!["https://example.test".to_string()],
941 allow_remote: true,
942 public_url: Some("wss://bridge.example.test/webmcp".to_string()),
943 ..Default::default()
944 };
945 let temp = TempDir::new().expect("temp dir");
946 let adapter = Arc::new(FilesystemWorkspace::new(temp.path(), [], false).await.expect("adapter"));
947 assert!(matches!(WebmcpServer::new(adapter, remote_config), Err(WebmcpError::InvalidRequest(_))));
948
949 let temp = TempDir::new().expect("temp dir");
950 let adapter = Arc::new(FilesystemWorkspace::new(temp.path(), [], false).await.expect("adapter"));
951 let valid_proxy_config = WebmcpServerConfig {
952 allowed_origins: vec!["https://example.test".to_string()],
953 allow_remote: true,
954 public_url: Some("wss://bridge.example.test/webmcp".to_string()),
955 ..Default::default()
956 };
957 assert!(WebmcpServer::new(adapter.clone(), valid_proxy_config).is_ok());
958 let invalid_public_url_config = WebmcpServerConfig {
959 allowed_origins: vec!["https://example.test".to_string()],
960 public_url: Some("ws://bridge.example.test/webmcp".to_string()),
961 ..Default::default()
962 };
963 assert!(matches!(WebmcpServer::new(adapter, invalid_public_url_config), Err(WebmcpError::InvalidRequest(_))));
964
965 let temp = TempDir::new().expect("temp dir");
966 let adapter = Arc::new(FilesystemWorkspace::new(temp.path(), [], false).await.expect("adapter"));
967 let invalid_public_url_config = WebmcpServerConfig {
968 allowed_origins: vec!["https://example.test".to_string()],
969 allow_remote: true,
970 public_url: Some("wss://".to_string()),
971 ..Default::default()
972 };
973 assert!(matches!(WebmcpServer::new(adapter, invalid_public_url_config), Err(WebmcpError::InvalidRequest(_))));
974 }
975}