Files
cursor-byok/server/src/cursor/proxy.rs
T

343 lines
10 KiB
Rust

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<reqwest::Client>,
store: Option<crate::store::Store>,
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<Body> {
let body = self.body.clone();
self.with_body(body)
}
pub fn with_body(mut self, body: Bytes) -> Response<Body> {
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<Self> {
Ok(Self {
client: None,
store: Some(store),
upstream: CURSOR_UPSTREAM.into(),
})
}
#[cfg(test)]
pub(crate) fn for_upstream(upstream: &str) -> Result<Self> {
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<reqwest::Client> {
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<CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
forward_request(&proxy, request, None).await
}
pub(crate) async fn forward_to_service(
proxy: &CursorProxy,
request: Request<Body>,
service_url: &str,
) -> Result<Response<Body>> {
forward_request(proxy, request, Some(service_url)).await
}
async fn forward_request(
proxy: &CursorProxy,
request: Request<Body>,
service_url: Option<&str>,
) -> Result<Response<Body>> {
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<Body>,
) -> Result<BufferedResponse> {
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<String> {
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::<Vec<_>>()
})
.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<Body>| 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<Body>| 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();
}
}