use std::time::Instant; use axum::{ body::{to_bytes, Body, Bytes}, extract::Extension, http::{header, Request, Response}, }; use crate::Result; const CURSOR_UPSTREAM: &str = "https://api2.cursor.sh"; pub const UPSTREAM_URL_HEADER: &str = "x-server-upstream-url"; #[derive(Clone)] pub struct CursorProxy { client: Option, store: Option, upstream: String, } pub struct BufferedResponse { pub status: axum::http::StatusCode, pub headers: axum::http::HeaderMap, pub body: Bytes, } impl BufferedResponse { pub fn into_response(self) -> Response { let body = self.body.clone(); self.with_body(body) } pub fn with_body(mut self, body: Bytes) -> Response { self.headers.insert( header::CONTENT_LENGTH, body.len() .to_string() .parse() .expect("body length is always a valid header value"), ); let mut response = Response::new(Body::from(body)); *response.status_mut() = self.status; *response.headers_mut() = self.headers; response } } impl CursorProxy { pub fn cursor(store: crate::store::Store) -> Result { Ok(Self { client: None, store: Some(store), upstream: CURSOR_UPSTREAM.into(), }) } #[cfg(test)] pub(crate) fn for_upstream(upstream: &str) -> Result { Ok(Self { client: Some( reqwest::Client::builder() .redirect(reqwest::redirect::Policy::none()) .build()?, ), store: None, upstream: upstream.trim_end_matches('/').to_owned(), }) } async fn client(&self) -> Result { match (&self.client, &self.store) { (Some(client), _) => Ok(client.clone()), (_, Some(store)) => Ok(crate::network::client_builder(store) .await? .redirect(reqwest::redirect::Policy::none()) .build()?), _ => unreachable!("Cursor proxy always has a client or store"), } } } pub async fn forward( Extension(proxy): Extension, request: Request, ) -> Result> { forward_request(&proxy, request, None).await } pub(crate) async fn forward_to_service( proxy: &CursorProxy, request: Request, service_url: &str, ) -> Result> { forward_request(proxy, request, Some(service_url)).await } async fn forward_request( proxy: &CursorProxy, request: Request, service_url: Option<&str>, ) -> Result> { let started = Instant::now(); let (parts, body) = request.into_parts(); let path = parts .uri .path_and_query() .map_or("/", |value| value.as_str()) .to_owned(); let url = match service_url { Some(service_url) => format!("{}{}", service_url.trim_end_matches('/'), path), None => upstream_url(&parts.headers, &proxy.upstream, &path)?, }; let mut headers = parts.headers; headers.remove(UPSTREAM_URL_HEADER); headers.remove(header::HOST); remove_hop_by_hop_headers(&mut headers); let client = proxy.client().await?; let upstream = client .request(parts.method.clone(), url) .headers(headers) .body(reqwest::Body::wrap_stream(body.into_data_stream())) .send() .await; let upstream = match upstream { Ok(response) => response, Err(error) => { tracing::error!( method = %parts.method, path, elapsed_ms = started.elapsed().as_millis(), %error, "Cursor upstream request failed" ); return Err(error.into()); } }; let status = upstream.status(); let mut response_headers = upstream.headers().clone(); remove_hop_by_hop_headers(&mut response_headers); let mut response = Response::new(Body::from_stream(upstream.bytes_stream())); *response.status_mut() = status; *response.headers_mut() = response_headers; tracing::info!( method = %parts.method, path, %status, elapsed_ms = started.elapsed().as_millis(), "forwarded Cursor backend request" ); Ok(response) } pub async fn forward_buffered( proxy: &CursorProxy, request: Request, ) -> Result { let (parts, body) = request.into_parts(); let path = parts .uri .path_and_query() .map_or("/", |value| value.as_str()); let url = upstream_url(&parts.headers, &proxy.upstream, path)?; let mut headers = parts.headers; headers.remove(UPSTREAM_URL_HEADER); headers.remove(header::HOST); remove_hop_by_hop_headers(&mut headers); headers.insert( "connect-accept-encoding", axum::http::HeaderValue::from_static("identity"), ); headers.insert( header::ACCEPT_ENCODING, axum::http::HeaderValue::from_static("identity"), ); let body = to_bytes(body, usize::MAX) .await .map_err(|error| crate::Error::Protocol(format!("cannot read request body: {error}")))?; let upstream = proxy .client() .await? .request(parts.method, url) .headers(headers) .body(body) .send() .await?; let status = upstream.status(); let mut headers = upstream.headers().clone(); remove_hop_by_hop_headers(&mut headers); let body = upstream.bytes().await?; Ok(BufferedResponse { status, headers, body, }) } fn upstream_url(headers: &axum::http::HeaderMap, fallback: &str, path: &str) -> Result { let Some(value) = headers.get(UPSTREAM_URL_HEADER) else { return Ok(format!("{fallback}{path}")); }; let value = value .to_str() .map_err(|error| crate::Error::Protocol(format!("invalid upstream URL header: {error}")))?; let url = reqwest::Url::parse(value) .map_err(|error| crate::Error::Protocol(format!("invalid upstream URL: {error}")))?; let host = url.host_str().unwrap_or_default(); if url.scheme() != "https" || !crate::harness::proxy_host_allowed(host) { return Err(crate::Error::Protocol( "upstream URL must target a Cursor HTTPS host".into(), )); } Ok(url.into()) } fn remove_hop_by_hop_headers(headers: &mut axum::http::HeaderMap) { let connection_headers = headers .get(header::CONNECTION) .and_then(|value| value.to_str().ok()) .map(|value| { value .split(',') .map(str::trim) .filter(|name| !name.is_empty()) .map(str::to_owned) .collect::>() }) .unwrap_or_default(); for name in connection_headers { headers.remove(name); } for name in [ header::CONNECTION, header::PROXY_AUTHENTICATE, header::PROXY_AUTHORIZATION, header::TE, header::TRAILER, header::TRANSFER_ENCODING, header::UPGRADE, ] { headers.remove(name); } headers.remove("keep-alive"); } #[cfg(test)] mod tests { use axum::{ body::{to_bytes, Body}, extract::Extension, http::{header, Request, StatusCode}, response::IntoResponse, routing::any, Router, }; use tower::ServiceExt; use super::{forward, forward_to_service, CursorProxy}; #[tokio::test] async fn preserves_request_and_response() { let upstream = Router::new().route( "/unknown", any(|request: Request| async move { let method = request.method().clone(); let query = request.uri().query().unwrap_or_default().to_owned(); let marker = request.headers()["x-marker"].clone(); let body = to_bytes(request.into_body(), usize::MAX).await.unwrap(); ( StatusCode::CREATED, [(header::CONTENT_TYPE, "application/proto")], format!( "{method} {query} {} {}", marker.to_str().unwrap(), String::from_utf8_lossy(&body) ), ) .into_response() }), ); let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); let address = listener.local_addr().unwrap(); let server = tokio::spawn(async move { axum::serve(listener, upstream).await.unwrap() }); let proxy = CursorProxy::for_upstream(&format!("http://{address}")).unwrap(); let app = Router::new().fallback(forward).layer(Extension(proxy)); let response = app .oneshot( Request::put("/unknown?a=1") .header("x-marker", "kept") .body(Body::from("payload")) .unwrap(), ) .await .unwrap(); assert_eq!(response.status(), StatusCode::CREATED); assert_eq!( response.headers()[header::CONTENT_TYPE], "application/proto" ); assert_eq!( to_bytes(response.into_body(), usize::MAX).await.unwrap(), "PUT a=1 kept payload" ); server.abort(); } #[tokio::test] async fn tab_service_keeps_its_base_path_and_the_original_query() { let upstream = Router::new().route( "/base/aiserver.v1.AiService/StreamCpp", any(|request: Request| async move { request.uri().path_and_query().unwrap().as_str().to_owned() }), ); let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); let address = listener.local_addr().unwrap(); let server = tokio::spawn(async move { axum::serve(listener, upstream).await.unwrap() }); let proxy = CursorProxy::for_upstream("http://unused.invalid").unwrap(); let response = forward_to_service( &proxy, Request::post("/aiserver.v1.AiService/StreamCpp?client=cursor") .body(Body::empty()) .unwrap(), &format!("http://{address}/base"), ) .await .unwrap(); assert_eq!( to_bytes(response.into_body(), usize::MAX).await.unwrap(), "/base/aiserver.v1.AiService/StreamCpp?client=cursor" ); server.abort(); } }