1use std::{
2 collections::HashSet,
3 convert::Infallible,
4 pin::Pin,
5 sync::{Arc, Mutex as StdMutex},
6 task::{Context, Poll},
7};
8
9use axum::{
10 extract::{rejection::JsonRejection, DefaultBodyLimit, Path, Query, State},
11 http::StatusCode,
12 response::{
13 sse::{Event as SseEvent, KeepAlive, Sse},
14 IntoResponse, Json, Response,
15 },
16 routing::{delete, get, post},
17 Router,
18};
19use futures_core::Stream;
20use serde::Deserialize;
21use tokio::sync::{OwnedSemaphorePermit, Semaphore};
22
23use crate::Error;
24
25use super::{
26 backend::CodexRestBackend,
27 types::{
28 RestApprovalPolicy, RestBackend, RestCallBody, RestCallRequest, RestClientOptions,
29 RestError, RestErrorReplyRequest, RestErrorResponse, RestEventResponse, RestFuture,
30 RestHealthResponse, RestListSessionsResponse, RestRequestReplyResultRequest, RestResult,
31 RestRouterOptions, RestSessionCreateRequest, RestTextTurnRequest,
32 },
33};
34
35#[derive(Clone)]
36struct RestState {
37 backend: Arc<dyn RestBackend>,
38 options: RestRouterOptions,
39 one_shot_gate: Arc<Semaphore>,
40 active_polls: Arc<StdMutex<HashSet<String>>>,
41}
42
43struct ActivePollGuard {
44 active_polls: Arc<StdMutex<HashSet<String>>>,
45 session_id: String,
46}
47
48impl Drop for ActivePollGuard {
49 fn drop(&mut self) {
50 let mut active = self
51 .active_polls
52 .lock()
53 .unwrap_or_else(|poisoned| poisoned.into_inner());
54 active.remove(&self.session_id);
55 }
56}
57
58pub fn router() -> Router {
65 router_with_options(RestRouterOptions::default())
66}
67
68pub fn text_turn_router() -> Router {
70 router_with_options(RestRouterOptions::text_turn())
71}
72
73pub fn trusted_bridge_router() -> Router {
91 router_with_options(RestRouterOptions::trusted_bridge())
92}
93
94pub fn router_with_options(options: RestRouterOptions) -> Router {
96 router_with_backend_and_options(
97 CodexRestBackend::with_limits(options.limits.clone()),
98 options,
99 )
100}
101
102pub fn router_with_backend<B>(backend: B) -> Router
104where
105 B: RestBackend,
106{
107 router_with_backend_and_options(backend, RestRouterOptions::default())
108}
109
110pub fn router_with_backend_and_options<B>(backend: B, options: RestRouterOptions) -> Router
112where
113 B: RestBackend,
114{
115 router_with_backend_arc_and_options(Arc::new(backend), options)
116}
117
118pub fn router_with_backend_arc(backend: Arc<dyn RestBackend>) -> Router {
120 router_with_backend_arc_and_options(backend, RestRouterOptions::default())
121}
122
123pub fn router_with_backend_arc_and_options(
125 backend: Arc<dyn RestBackend>,
126 options: RestRouterOptions,
127) -> Router {
128 let state = RestState {
129 backend,
130 one_shot_gate: Arc::new(Semaphore::new(options.limits.max_one_shot_concurrency)),
131 active_polls: Arc::default(),
132 options: options.clone(),
133 };
134
135 let router = Router::new()
136 .route("/health", get(health))
137 .route("/v1/health", get(health))
138 .route("/v1/compatibility", get(compatibility));
139
140 let router = if options.enable_text_turn_route {
141 router.route("/v1/text-turn", post(text_turn))
142 } else {
143 router
144 };
145
146 let router = if options.enable_bridge_routes {
147 router
148 .route("/v1/call/{*method}", post(call_method))
149 .route("/v1/sessions", get(list_sessions).post(create_session))
150 .route("/v1/sessions/{session_id}", delete(delete_session))
151 .route(
152 "/v1/sessions/{session_id}/call/{*method}",
153 post(call_session_method),
154 )
155 .route("/v1/sessions/{session_id}/events", get(poll_event))
156 .route(
157 "/v1/sessions/{session_id}/events/stream",
158 get(poll_event_stream),
159 )
160 .route(
161 "/v1/sessions/{session_id}/requests/{request_key}/result",
162 post(reply_request_result),
163 )
164 .route(
165 "/v1/sessions/{session_id}/requests/{request_key}/error",
166 post(reply_request_error),
167 )
168 } else {
169 router
170 };
171
172 router
179 .layer(DefaultBodyLimit::max(options.limits.max_request_body_bytes))
180 .with_state(state)
181}
182
183#[derive(Clone, Debug, Deserialize)]
184#[serde(rename_all = "camelCase")]
185struct EventQuery {
186 timeout_ms: Option<u64>,
187}
188
189async fn health() -> impl IntoResponse {
190 Json(RestHealthResponse {
191 status: "ok".to_owned(),
192 })
193}
194
195async fn compatibility(State(state): State<RestState>) -> impl IntoResponse {
196 match state.backend.compatibility_report().await {
197 Ok(response) => Json(response).into_response(),
198 Err(error) => rest_error(error),
199 }
200}
201
202async fn text_turn(
203 State(state): State<RestState>,
204 body: std::result::Result<Json<RestTextTurnRequest>, JsonRejection>,
205) -> Response {
206 let Json(request) = match body {
207 Ok(body) => body,
208 Err(error) => return invalid_json(error),
209 };
210 if request.prompt.trim().is_empty() {
211 return invalid_request("prompt must not be empty");
212 }
213 if let Err(error) = validate_text_turn_request(&state.options, &request) {
214 return rest_error(error);
215 }
216
217 let _permit = match acquire_one_shot_permit(&state) {
218 Ok(permit) => permit,
219 Err(error) => return rest_error(error),
220 };
221 match state.backend.run_text_turn(request).await {
222 Ok(response) => Json(response).into_response(),
223 Err(error) => rest_error(error),
224 }
225}
226
227async fn call_method(
228 State(state): State<RestState>,
229 Path(method): Path<String>,
230 body: std::result::Result<Json<RestCallBody>, JsonRejection>,
231) -> Response {
232 let method = match normalize_method(method) {
233 Some(method) => method,
234 None => return invalid_request("method path must not be empty"),
235 };
236 let Json(body) = match body {
237 Ok(body) => body,
238 Err(error) => return invalid_json(error),
239 };
240 if let Err(error) = validate_client_options(&state.options, body.client.as_ref()) {
241 return rest_error(error);
242 }
243 let _permit = match acquire_one_shot_permit(&state) {
244 Ok(permit) => permit,
245 Err(error) => return rest_error(error),
246 };
247 let request = RestCallRequest {
248 session_id: None,
249 method,
250 params: body.params,
251 client: body.client,
252 };
253 match state.backend.call_method(request).await {
254 Ok(response) => Json(response).into_response(),
255 Err(error) => rest_error(error),
256 }
257}
258
259async fn create_session(
260 State(state): State<RestState>,
261 body: std::result::Result<Json<RestSessionCreateRequest>, JsonRejection>,
262) -> Response {
263 let Json(request) = match body {
264 Ok(body) => body,
265 Err(error) => return invalid_json(error),
266 };
267 if let Err(error) = validate_client_options(&state.options, request.client.as_ref()) {
268 return rest_error(error);
269 }
270 match state.backend.list_sessions().await {
271 Ok(sessions) if sessions.len() >= state.options.limits.max_sessions => {
272 return rest_error(RestError::RateLimited(format!(
273 "maximum REST session count ({}) reached",
274 state.options.limits.max_sessions
275 )));
276 }
277 Ok(_) => {}
278 Err(error) => return rest_error(error),
279 }
280 match state.backend.create_session(request).await {
281 Ok(response) => Json(response).into_response(),
282 Err(error) => rest_error(error),
283 }
284}
285
286async fn list_sessions(State(state): State<RestState>) -> Response {
287 match state.backend.list_sessions().await {
288 Ok(sessions) => Json(RestListSessionsResponse { sessions }).into_response(),
289 Err(error) => rest_error(error),
290 }
291}
292
293async fn delete_session(
294 State(state): State<RestState>,
295 Path(session_id): Path<String>,
296) -> Response {
297 match state.backend.delete_session(session_id).await {
298 Ok(response) => Json(response).into_response(),
299 Err(error) => rest_error(error),
300 }
301}
302
303async fn call_session_method(
304 State(state): State<RestState>,
305 Path((session_id, method)): Path<(String, String)>,
306 body: std::result::Result<Json<RestCallBody>, JsonRejection>,
307) -> Response {
308 let method = match normalize_method(method) {
309 Some(method) => method,
310 None => return invalid_request("method path must not be empty"),
311 };
312 let Json(body) = match body {
313 Ok(body) => body,
314 Err(error) => return invalid_json(error),
315 };
316 if body.client.is_some() {
317 return rest_error(RestError::InvalidRequest(
318 "`client` options are only accepted when creating a session or making one-shot calls"
319 .to_owned(),
320 ));
321 }
322 let request = RestCallRequest {
323 session_id: Some(session_id),
324 method,
325 params: body.params,
326 client: body.client,
327 };
328 match state.backend.call_method(request).await {
329 Ok(response) => Json(response).into_response(),
330 Err(error) => rest_error(error),
331 }
332}
333
334async fn poll_event(
335 State(state): State<RestState>,
336 Path(session_id): Path<String>,
337 Query(query): Query<EventQuery>,
338) -> Response {
339 let _guard = match acquire_poll_guard(&state, &session_id) {
340 Ok(guard) => guard,
341 Err(error) => return rest_error(error),
342 };
343 let timeout_ms = Some(clamp_poll_timeout_ms(
344 &state.options,
345 query
346 .timeout_ms
347 .unwrap_or(state.options.limits.max_poll_timeout.as_millis() as u64),
348 ));
349 match state.backend.poll_event(session_id, timeout_ms).await {
350 Ok(response) => Json(response).into_response(),
351 Err(error) => rest_error(error),
352 }
353}
354
355async fn poll_event_stream(
387 State(state): State<RestState>,
388 Path(session_id): Path<String>,
389 Query(query): Query<EventQuery>,
390) -> Response {
391 let guard = match acquire_poll_guard(&state, &session_id) {
392 Ok(guard) => guard,
393 Err(error) => return rest_error(error),
394 };
395 let timeout_ms = clamp_stream_poll_timeout_ms(
396 &state.options,
397 query
398 .timeout_ms
399 .unwrap_or(state.options.limits.max_poll_timeout.as_millis() as u64),
400 );
401 let stream = EventPollStream {
402 backend: state.backend.clone(),
403 session_id,
404 timeout_ms,
405 pending: None,
406 guard: Some(guard),
407 done: false,
408 synchronous_polls: 0,
409 };
410 Sse::new(stream)
411 .keep_alive(KeepAlive::new().interval(state.options.limits.sse_keep_alive_interval))
412 .into_response()
413}
414
415struct EventPollStream {
428 backend: Arc<dyn RestBackend>,
429 session_id: String,
430 timeout_ms: u64,
431 pending: Option<RestFuture<RestEventResponse>>,
435 guard: Option<ActivePollGuard>,
441 done: bool,
442 synchronous_polls: u32,
447}
448
449const YIELD_AFTER_SYNCHRONOUS_POLLS: u32 = 32;
473
474impl Stream for EventPollStream {
475 type Item = Result<SseEvent, Infallible>;
476
477 fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
478 let this = self.get_mut();
479 if this.done {
480 return Poll::Ready(None);
481 }
482 if this.synchronous_polls >= YIELD_AFTER_SYNCHRONOUS_POLLS {
483 this.synchronous_polls = 0;
488 cx.waker().wake_by_ref();
489 return Poll::Pending;
490 }
491 if this.pending.is_none() {
492 this.pending = Some(
493 this.backend
494 .poll_event(this.session_id.clone(), Some(this.timeout_ms)),
495 );
496 }
497 let pending = this
498 .pending
499 .as_mut()
500 .expect("pending future was just populated above");
501 match pending.as_mut().poll(cx) {
502 Poll::Pending => {
503 this.synchronous_polls = 0;
506 Poll::Pending
507 }
508 Poll::Ready(result) => {
509 this.pending = None;
510 this.synchronous_polls = this.synchronous_polls.saturating_add(1);
511 match result {
512 Ok(response) => {
513 if matches!(response, RestEventResponse::Closed) {
514 this.done = true;
515 this.guard = None;
516 }
517 Poll::Ready(Some(Ok(sse_event_from_response(&response))))
518 }
519 Err(error) => {
520 this.done = true;
521 this.guard = None;
522 Poll::Ready(Some(Ok(sse_error_event(error))))
523 }
524 }
525 }
526 }
527 }
528}
529
530fn sse_event_from_response(response: &RestEventResponse) -> SseEvent {
535 let event_name = match response {
536 RestEventResponse::Notification { .. } => "notification",
537 RestEventResponse::Request { .. } => "request",
538 RestEventResponse::Closed => "closed",
539 RestEventResponse::Timeout => "timeout",
540 };
541 let payload = serde_json::to_string(response)
542 .unwrap_or_else(|_| r#"{"event":"internal_error"}"#.to_owned());
543 SseEvent::default().event(event_name).data(payload)
544}
545
546fn sse_error_event(error: RestError) -> SseEvent {
551 let (_status, body) = rest_error_response(error);
552 let payload =
553 serde_json::to_string(&body).unwrap_or_else(|_| r#"{"error":"internal"}"#.to_owned());
554 SseEvent::default().event("error").data(payload)
555}
556
557async fn reply_request_result(
558 State(state): State<RestState>,
559 Path((session_id, request_key)): Path<(String, String)>,
560 body: std::result::Result<Json<RestRequestReplyResultRequest>, JsonRejection>,
561) -> Response {
562 let Json(body) = match body {
563 Ok(body) => body,
564 Err(error) => return invalid_json(error),
565 };
566 match state
567 .backend
568 .reply_request_result(session_id, request_key, body)
569 .await
570 {
571 Ok(response) => Json(response).into_response(),
572 Err(error) => rest_error(error),
573 }
574}
575
576async fn reply_request_error(
577 State(state): State<RestState>,
578 Path((session_id, request_key)): Path<(String, String)>,
579 body: std::result::Result<Json<RestErrorReplyRequest>, JsonRejection>,
580) -> Response {
581 let Json(body) = match body {
582 Ok(body) => body,
583 Err(error) => return invalid_json(error),
584 };
585 match state
586 .backend
587 .reply_request_error(session_id, request_key, body)
588 .await
589 {
590 Ok(response) => Json(response).into_response(),
591 Err(error) => rest_error(error),
592 }
593}
594
595fn invalid_request(message: impl Into<String>) -> Response {
596 (
597 StatusCode::BAD_REQUEST,
598 Json(RestErrorResponse {
599 error: "invalid_request".to_owned(),
600 message: message.into(),
601 code: None,
602 data: None,
603 }),
604 )
605 .into_response()
606}
607
608fn invalid_json(error: JsonRejection) -> Response {
609 if error.status() == StatusCode::PAYLOAD_TOO_LARGE {
617 return (
618 StatusCode::PAYLOAD_TOO_LARGE,
619 Json(RestErrorResponse {
620 error: "payload_too_large".to_owned(),
621 message: error.body_text(),
622 code: None,
623 data: None,
624 }),
625 )
626 .into_response();
627 }
628 (
629 StatusCode::BAD_REQUEST,
630 Json(RestErrorResponse {
631 error: "invalid_json".to_owned(),
632 message: error.body_text(),
633 code: None,
634 data: None,
635 }),
636 )
637 .into_response()
638}
639
640fn rest_error_response(error: RestError) -> (StatusCode, RestErrorResponse) {
647 fn simple(status: StatusCode, kind: &str, message: String) -> (StatusCode, RestErrorResponse) {
653 (
654 status,
655 RestErrorResponse {
656 error: kind.to_owned(),
657 message,
658 code: None,
659 data: None,
660 },
661 )
662 }
663
664 match error {
665 RestError::NotFound(message) => simple(StatusCode::NOT_FOUND, "not_found", message),
666 RestError::Gone(message) => simple(StatusCode::GONE, "gone", message),
667 RestError::Forbidden(message) => simple(StatusCode::FORBIDDEN, "forbidden", message),
668 RestError::InvalidRequest(message) => {
669 simple(StatusCode::BAD_REQUEST, "invalid_request", message)
670 }
671 RestError::RateLimited(message) => {
672 simple(StatusCode::TOO_MANY_REQUESTS, "rate_limited", message)
673 }
674 RestError::Conflict(message) => simple(StatusCode::CONFLICT, "conflict", message),
675 RestError::TimedOut(message) => simple(StatusCode::GATEWAY_TIMEOUT, "timeout", message),
676 RestError::PayloadTooLarge(message) => {
677 simple(StatusCode::PAYLOAD_TOO_LARGE, "payload_too_large", message)
678 }
679 RestError::Internal(message) => {
680 simple(StatusCode::INTERNAL_SERVER_ERROR, "internal", message)
681 }
682 RestError::Client(Error::Rpc {
685 code,
686 message,
687 data,
688 }) => (
689 StatusCode::BAD_GATEWAY,
690 RestErrorResponse {
691 error: "json_rpc_error".to_owned(),
692 message,
693 code: Some(code),
694 data,
695 },
696 ),
697 RestError::Client(error) => simple(
698 StatusCode::BAD_GATEWAY,
699 "codex_app_server_error",
700 error.to_string(),
701 ),
702 }
703}
704
705fn rest_error(error: RestError) -> Response {
706 let (status, body) = rest_error_response(error);
707 (status, Json(body)).into_response()
708}
709
710fn validate_text_turn_request(
711 options: &RestRouterOptions,
712 request: &RestTextTurnRequest,
713) -> RestResult<()> {
714 if !options.allow_unsafe_client_options
715 && matches!(request.approval_policy, Some(RestApprovalPolicy::AllowAll))
716 {
717 return Err(RestError::Forbidden(
718 "`approvalPolicy: allow_all` requires a trusted REST bridge".to_owned(),
719 ));
720 }
721 validate_client_options(options, request.client.as_ref())
722}
723
724fn validate_client_options(
725 options: &RestRouterOptions,
726 client: Option<&RestClientOptions>,
727) -> RestResult<()> {
728 if options.allow_unsafe_client_options {
729 return Ok(());
730 }
731 let Some(client) = client else {
732 return Ok(());
733 };
734 if client.command.is_some() || !client.extra_args.is_empty() || !client.config.is_empty() {
735 return Err(RestError::Forbidden(
736 "client command, extraArgs, and config overrides require a trusted REST bridge"
737 .to_owned(),
738 ));
739 }
740 Ok(())
741}
742
743fn acquire_one_shot_permit(state: &RestState) -> RestResult<OwnedSemaphorePermit> {
744 state
745 .one_shot_gate
746 .clone()
747 .try_acquire_owned()
748 .map_err(|_| {
749 RestError::RateLimited(format!(
750 "maximum one-shot REST call concurrency ({}) reached",
751 state.options.limits.max_one_shot_concurrency
752 ))
753 })
754}
755
756fn acquire_poll_guard(state: &RestState, session_id: &str) -> RestResult<ActivePollGuard> {
757 let mut active = state
758 .active_polls
759 .lock()
760 .unwrap_or_else(|poisoned| poisoned.into_inner());
761 if !active.insert(session_id.to_owned()) {
762 return Err(RestError::Conflict(format!(
763 "an event poll is already active for session `{session_id}`"
764 )));
765 }
766 Ok(ActivePollGuard {
767 active_polls: state.active_polls.clone(),
768 session_id: session_id.to_owned(),
769 })
770}
771
772fn clamp_poll_timeout_ms(options: &RestRouterOptions, timeout_ms: u64) -> u64 {
773 let max = options.limits.max_poll_timeout.as_millis() as u64;
774 timeout_ms.min(max)
775}
776
777fn clamp_stream_poll_timeout_ms(options: &RestRouterOptions, timeout_ms: u64) -> u64 {
797 let min = options.limits.min_stream_poll_timeout.as_millis() as u64;
798 let max = options.limits.max_poll_timeout.as_millis() as u64;
799 timeout_ms.min(max).max(min)
803}
804
805fn normalize_method(method: String) -> Option<String> {
806 let method = method.trim_matches('/').trim();
807 (!method.is_empty()).then(|| method.to_owned())
808}