mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-06 21:52:51 +08:00
feat: initialize server structure and database schema
- Added initial server setup with Cargo.toml defining dependencies and project structure. - Created build.rs for generating protobuf bindings and validating wire contracts. - Established database schema with initial migration files for conversations, messages, and runs. - Introduced tools and prompts for Cursor functionality, enhancing user interaction capabilities.
This commit is contained in:
@@ -0,0 +1,207 @@
|
||||
//! Accepts ordered Cursor Bidi append requests and routes them by request_id.
|
||||
use prost::Message;
|
||||
|
||||
use crate::{
|
||||
cursor::{
|
||||
conversation::TransportCommand,
|
||||
protocol::{
|
||||
events,
|
||||
proto::{agent::v1 as agent, aiserver::v1 as ai},
|
||||
},
|
||||
transport::{TransportParent, TransportRegistry},
|
||||
},
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
pub struct DecodedAppend {
|
||||
pub request_id: String,
|
||||
pub seqno: i64,
|
||||
pub message: agent::AgentClientMessage,
|
||||
}
|
||||
|
||||
impl DecodedAppend {
|
||||
pub fn model_id(&self) -> Option<&str> {
|
||||
let agent::agent_client_message::Message::RunRequest(request) =
|
||||
self.message.message.as_ref()?
|
||||
else {
|
||||
return None;
|
||||
};
|
||||
request
|
||||
.requested_model
|
||||
.as_ref()
|
||||
.map(|model| model.model_id.as_str())
|
||||
.filter(|model| !model.is_empty())
|
||||
.or_else(|| {
|
||||
request
|
||||
.model_details
|
||||
.as_ref()
|
||||
.map(|model| model.model_id.as_str())
|
||||
.filter(|model| !model.is_empty())
|
||||
})
|
||||
}
|
||||
|
||||
pub fn conversation_id(&self) -> Option<&str> {
|
||||
let agent::agent_client_message::Message::RunRequest(request) =
|
||||
self.message.message.as_ref()?
|
||||
else {
|
||||
return None;
|
||||
};
|
||||
request.conversation_id.as_deref()
|
||||
}
|
||||
|
||||
pub fn is_background_task_completion(&self) -> bool {
|
||||
let Some(agent::agent_client_message::Message::RunRequest(request)) =
|
||||
self.message.message.as_ref()
|
||||
else {
|
||||
return false;
|
||||
};
|
||||
matches!(
|
||||
request
|
||||
.action
|
||||
.as_ref()
|
||||
.and_then(|action| action.action.as_ref()),
|
||||
Some(agent::conversation_action::Action::BackgroundTaskCompletionAction(_))
|
||||
)
|
||||
}
|
||||
|
||||
pub fn trace_metadata(&self) -> serde_json::Value {
|
||||
let Some(message) = self.message.message.as_ref() else {
|
||||
return serde_json::json!({
|
||||
"append_seqno": self.seqno,
|
||||
"message_type": "empty",
|
||||
});
|
||||
};
|
||||
let agent::agent_client_message::Message::RunRequest(request) = message else {
|
||||
return serde_json::json!({
|
||||
"append_seqno": self.seqno,
|
||||
"message_type": client_message_type(message),
|
||||
});
|
||||
};
|
||||
let (action_type, history_messages, history_images) = request
|
||||
.action
|
||||
.as_ref()
|
||||
.and_then(|action| action.action.as_ref())
|
||||
.map(|action| match action {
|
||||
agent::conversation_action::Action::UserMessageAction(action) => {
|
||||
let history = action.conversation_history.as_ref();
|
||||
(
|
||||
"user_message",
|
||||
history.map_or(0, |history| history.messages.len()),
|
||||
history.map_or(0, history_image_count),
|
||||
)
|
||||
}
|
||||
agent::conversation_action::Action::BackgroundTaskCompletionAction(_) => {
|
||||
("background_task_completion", 0, 0)
|
||||
}
|
||||
agent::conversation_action::Action::ExecutePlanAction(_) => ("execute_plan", 0, 0),
|
||||
agent::conversation_action::Action::SummarizeAction(_) => ("summarize", 0, 0),
|
||||
_ => ("other", 0, 0),
|
||||
})
|
||||
.unwrap_or(("none", 0, 0));
|
||||
let state = request.conversation_state.as_ref();
|
||||
serde_json::json!({
|
||||
"append_seqno": self.seqno,
|
||||
"message_type": "run_request",
|
||||
"conversation_id": request.conversation_id,
|
||||
"model_id": self.model_id(),
|
||||
"action_type": action_type,
|
||||
"conversation_history_messages": history_messages,
|
||||
"conversation_history_images": history_images,
|
||||
"root_message_count": state.map_or(0, |state| state.root_prompt_messages_json.len()),
|
||||
"turn_count": state.map_or(0, |state| state.turns.len()),
|
||||
"prefetched_blob_count": request.pre_fetched_blobs.len(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn client_message_type(message: &agent::agent_client_message::Message) -> &'static str {
|
||||
use agent::agent_client_message::Message;
|
||||
match message {
|
||||
Message::RunRequest(_) => "run_request",
|
||||
Message::ExecClientMessage(_) => "exec_client_message",
|
||||
Message::ExecClientControlMessage(_) => "exec_client_control_message",
|
||||
Message::KvClientMessage(_) => "kv_client_message",
|
||||
Message::ConversationAction(_) => "conversation_action",
|
||||
Message::InteractionResponse(_) => "interaction_response",
|
||||
Message::ClientHeartbeat(_) => "client_heartbeat",
|
||||
Message::PrewarmRequest(_) => "prewarm_request",
|
||||
}
|
||||
}
|
||||
|
||||
fn history_image_count(history: &agent::ConversationHistory) -> usize {
|
||||
use agent::{
|
||||
conversation_history_message::Message,
|
||||
conversation_history_tool_result_content::Content as ToolContent,
|
||||
conversation_history_user_content::Content as UserContent,
|
||||
};
|
||||
history
|
||||
.messages
|
||||
.iter()
|
||||
.map(|message| match message.message.as_ref() {
|
||||
Some(Message::User(user)) => user
|
||||
.content
|
||||
.iter()
|
||||
.filter(|content| matches!(content.content, Some(UserContent::Image(_))))
|
||||
.count(),
|
||||
Some(Message::Tool(tool)) => tool
|
||||
.content
|
||||
.iter()
|
||||
.filter(|content| matches!(content.content, Some(ToolContent::Image(_))))
|
||||
.count(),
|
||||
_ => 0,
|
||||
})
|
||||
.sum()
|
||||
}
|
||||
|
||||
pub fn decode(request: &ai::BidiAppendRequest) -> Result<DecodedAppend> {
|
||||
let request_id = request
|
||||
.request_id
|
||||
.as_ref()
|
||||
.map(|id| id.request_id.as_str())
|
||||
.filter(|id| !id.is_empty())
|
||||
.ok_or_else(|| Error::Protocol("BidiAppend request_id is required".into()))?;
|
||||
if !request.data_binary.is_empty() {
|
||||
return Err(Error::Protocol(
|
||||
"BidiAppend data_binary is not part of the captured protocol".into(),
|
||||
));
|
||||
}
|
||||
if request.data.is_empty() {
|
||||
return Err(Error::Protocol(
|
||||
"BidiAppend contains no AgentClientMessage".into(),
|
||||
));
|
||||
}
|
||||
let payload = hex::decode(&request.data)
|
||||
.map_err(|error| Error::Protocol(format!("invalid BidiAppend hex: {error}")))?;
|
||||
Ok(DecodedAppend {
|
||||
request_id: request_id.into(),
|
||||
seqno: request.append_seqno,
|
||||
message: agent::AgentClientMessage::decode(payload.as_slice())?,
|
||||
})
|
||||
}
|
||||
|
||||
pub async fn append(
|
||||
registry: &TransportRegistry,
|
||||
request: DecodedAppend,
|
||||
parent: Option<TransportParent>,
|
||||
) -> Result<ai::BidiAppendResponse> {
|
||||
let handle = registry.get_or_create(&request.request_id).await?;
|
||||
if let Some(conversation_id) = request.conversation_id() {
|
||||
handle.set_conversation_id(conversation_id)?;
|
||||
}
|
||||
if let Some(parent) = parent {
|
||||
handle.set_parent(parent)?;
|
||||
}
|
||||
if matches!(
|
||||
request.message.message.as_ref(),
|
||||
Some(agent::agent_client_message::Message::ClientHeartbeat(_))
|
||||
) {
|
||||
handle.emit(&events::heartbeat())?;
|
||||
}
|
||||
handle
|
||||
.command(TransportCommand::Append {
|
||||
seqno: request.seqno,
|
||||
message: Box::new(request.message),
|
||||
})
|
||||
.await?;
|
||||
Ok(ai::BidiAppendResponse {})
|
||||
}
|
||||
@@ -0,0 +1,227 @@
|
||||
//! Implements Cursor HTTP endpoints outside the Agent Run stream.
|
||||
use axum::{
|
||||
body::{to_bytes, Body, Bytes},
|
||||
extract::{DefaultBodyLimit, Extension, State},
|
||||
http::{header, HeaderMap, HeaderValue, Request, Response, StatusCode},
|
||||
routing::{get, post},
|
||||
Router,
|
||||
};
|
||||
use tower_http::decompression::RequestDecompressionLayer;
|
||||
|
||||
use crate::{
|
||||
api::cursor::{
|
||||
bidi,
|
||||
proxy::{self, CursorProxy},
|
||||
run_sse,
|
||||
},
|
||||
cursor::{
|
||||
protocol::{
|
||||
connect,
|
||||
proto::{agent::v1 as agent, aiserver::v1 as ai},
|
||||
},
|
||||
services::{account, analytics, model_catalog, observability::CursorTraceRecorder, tab},
|
||||
transport::{TransportParent, TransportRegistry},
|
||||
},
|
||||
Result,
|
||||
};
|
||||
|
||||
pub fn router(registry: TransportRegistry) -> Result<Router> {
|
||||
let proxy = CursorProxy::cursor(registry.store().clone())?;
|
||||
Ok(router_with_proxy(registry, proxy))
|
||||
}
|
||||
|
||||
fn router_with_proxy(registry: TransportRegistry, proxy: CursorProxy) -> Router {
|
||||
Router::new()
|
||||
.route("/__byok-api__/healthz", get(health))
|
||||
.route("/agent.v1.AgentService/RunSSE", post(run_sse_handler))
|
||||
.route("/aiserver.v1.BidiService/BidiAppend", post(bidi_handler))
|
||||
.route(
|
||||
"/aiserver.v1.AiService/AvailableModels",
|
||||
post(model_catalog::available_models),
|
||||
)
|
||||
.route(
|
||||
"/agent.v1.AgentService/GetUsableModels",
|
||||
post(model_catalog::usable_models),
|
||||
)
|
||||
.route(
|
||||
"/aiserver.v1.AiService/GetUsableModels",
|
||||
post(model_catalog::usable_models),
|
||||
)
|
||||
.route(
|
||||
"/aiserver.v1.AuthService/GetEmail",
|
||||
post(account::get_email),
|
||||
)
|
||||
.route("/aiserver.v1.DashboardService/GetMe", post(account::get_me))
|
||||
.route(
|
||||
"/aiserver.v1.DashboardService/GetTeams",
|
||||
post(account::get_teams),
|
||||
)
|
||||
.route(
|
||||
"/aiserver.v1.DashboardService/GetUserProfile",
|
||||
post(account::get_user_profile),
|
||||
)
|
||||
.route(
|
||||
"/aiserver.v1.DashboardService/GetCurrentPeriodUsage",
|
||||
post(account::current_period_usage),
|
||||
)
|
||||
.route(
|
||||
"/aiserver.v1.DashboardService/GetUsageLimitStatusAndActiveGrants",
|
||||
post(account::usage_limit_status),
|
||||
)
|
||||
.route(
|
||||
analytics::BOOTSTRAP_STATSIG_PATH,
|
||||
post(analytics::bootstrap_statsig),
|
||||
)
|
||||
.route("/auth/full_stripe_profile", get(account::stripe_profile))
|
||||
.merge(tab::router())
|
||||
.route_layer(DefaultBodyLimit::disable())
|
||||
.route_layer(RequestDecompressionLayer::new())
|
||||
.fallback(proxy::forward)
|
||||
.method_not_allowed_fallback(proxy::forward)
|
||||
.layer(Extension(proxy))
|
||||
.with_state(registry)
|
||||
}
|
||||
|
||||
async fn health() -> StatusCode {
|
||||
StatusCode::NO_CONTENT
|
||||
}
|
||||
|
||||
async fn run_sse_handler(
|
||||
State(registry): State<TransportRegistry>,
|
||||
Extension(proxy): Extension<CursorProxy>,
|
||||
request: Request<Body>,
|
||||
) -> Result<Response<Body>> {
|
||||
let (parts, body) = buffered(request).await?;
|
||||
let request: agent::BidiRequestId = connect::decode_unary(&body)?;
|
||||
let route = registry.wait_route(&request.request_id).await;
|
||||
let trace = CursorTraceRecorder::resume(registry.store().clone(), &request.request_id).await;
|
||||
if let Some(trace) = &trace {
|
||||
trace
|
||||
.request(
|
||||
"run_sse_request",
|
||||
&body,
|
||||
serde_json::json!({"request_id": request.request_id}),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
match route {
|
||||
crate::cursor::transport::TransportRoute::Local => {
|
||||
run_sse::stream(®istry, &request.request_id).await
|
||||
}
|
||||
crate::cursor::transport::TransportRoute::Upstream(generation) => {
|
||||
let response = proxy::forward(
|
||||
Extension(proxy),
|
||||
Request::from_parts(parts, Body::from(body)),
|
||||
)
|
||||
.await?;
|
||||
Ok(run_sse::upstream(registry, request.request_id, generation, response, trace).await)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn bidi_handler(
|
||||
State(registry): State<TransportRegistry>,
|
||||
Extension(proxy): Extension<CursorProxy>,
|
||||
request: Request<Body>,
|
||||
) -> Result<Response<Body>> {
|
||||
let (parts, body) = buffered(request).await?;
|
||||
let request: ai::BidiAppendRequest = connect::decode_unary(&body)?;
|
||||
let decoded = bidi::decode(&request)?;
|
||||
let first_model = decoded.model_id().map(str::to_owned);
|
||||
let conversation_id = decoded.conversation_id().map(str::to_owned);
|
||||
let trace_metadata = decoded.trace_metadata();
|
||||
let local = if let Some(model_id) = decoded.model_id() {
|
||||
if registry.store().model(model_id).await?.is_some() {
|
||||
tracing::info!(
|
||||
request_id = decoded.request_id,
|
||||
model_id,
|
||||
"routing Cursor Run to BYOK provider"
|
||||
);
|
||||
true
|
||||
} else {
|
||||
tracing::info!(
|
||||
request_id = decoded.request_id,
|
||||
model_id,
|
||||
"routing Cursor Run to Cursor upstream"
|
||||
);
|
||||
false
|
||||
}
|
||||
} else if registry.local(&decoded.request_id).await.is_some() {
|
||||
true
|
||||
} else if registry.upstream(&decoded.request_id).await {
|
||||
false
|
||||
} else {
|
||||
return Err(crate::Error::Protocol(
|
||||
"first BidiAppend message must select a model".into(),
|
||||
));
|
||||
};
|
||||
let trace = if first_model.is_some() {
|
||||
CursorTraceRecorder::begin(
|
||||
registry.store().clone(),
|
||||
&decoded.request_id,
|
||||
conversation_id.as_deref(),
|
||||
if local {
|
||||
"local_byok"
|
||||
} else {
|
||||
"cursor_official"
|
||||
},
|
||||
first_model.as_deref(),
|
||||
)
|
||||
.await
|
||||
} else {
|
||||
CursorTraceRecorder::resume(registry.store().clone(), &decoded.request_id).await
|
||||
};
|
||||
if let Some(trace) = &trace {
|
||||
trace.request("bidi_request", &body, trace_metadata).await;
|
||||
}
|
||||
if !local {
|
||||
if first_model.is_some() {
|
||||
registry.mark_upstream(&decoded.request_id).await;
|
||||
}
|
||||
return proxy::forward(
|
||||
Extension(proxy),
|
||||
Request::from_parts(parts, Body::from(body)),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
let parent = parent_headers(&parts.headers)?;
|
||||
bidi::append(®istry, decoded, parent).await?;
|
||||
let mut response = Response::new(axum::body::Body::empty());
|
||||
*response.status_mut() = StatusCode::OK;
|
||||
response.headers_mut().insert(
|
||||
header::CONTENT_TYPE,
|
||||
HeaderValue::from_static("application/proto"),
|
||||
);
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
async fn buffered(request: Request<Body>) -> Result<(axum::http::request::Parts, Bytes)> {
|
||||
let (parts, body) = request.into_parts();
|
||||
let body = to_bytes(body, usize::MAX)
|
||||
.await
|
||||
.map_err(|error| crate::Error::Protocol(format!("cannot read request body: {error}")))?;
|
||||
Ok((parts, body))
|
||||
}
|
||||
|
||||
fn parent_headers(headers: &HeaderMap) -> Result<Option<TransportParent>> {
|
||||
let request_id = header_text(headers, "x-parent-request-id")?;
|
||||
let tool_call_id = header_text(headers, "x-parent-agent-tool-call-id")?;
|
||||
match (request_id, tool_call_id) {
|
||||
(None, None) => Ok(None),
|
||||
(Some(request_id), Some(tool_call_id)) => Ok(Some(TransportParent {
|
||||
request_id: request_id.into(),
|
||||
tool_call_id: tool_call_id.into(),
|
||||
})),
|
||||
_ => Err(crate::Error::Protocol(
|
||||
"Cursor subagent request must include both parent headers".into(),
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn header_text<'a>(headers: &'a HeaderMap, name: &str) -> Result<Option<&'a str>> {
|
||||
headers
|
||||
.get(name)
|
||||
.map(|value| value.to_str())
|
||||
.transpose()
|
||||
.map_err(|error| crate::Error::Protocol(format!("invalid {name} header: {error}")))
|
||||
}
|
||||
@@ -0,0 +1,8 @@
|
||||
//! Wires the Cursor-facing API routes.
|
||||
|
||||
pub mod bidi;
|
||||
mod handlers;
|
||||
pub mod proxy;
|
||||
mod run_sse;
|
||||
|
||||
pub use handlers::router;
|
||||
@@ -0,0 +1,236 @@
|
||||
//! Selects local handling or the configured official Cursor upstream.
|
||||
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(),
|
||||
})
|
||||
}
|
||||
|
||||
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::local_app::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");
|
||||
}
|
||||
@@ -0,0 +1,232 @@
|
||||
//! Subscribes Cursor RunSSE clients to replayable Transport output.
|
||||
use axum::{
|
||||
body::Body,
|
||||
http::{header, HeaderValue, Response, StatusCode},
|
||||
};
|
||||
use bytes::Bytes;
|
||||
use std::convert::Infallible;
|
||||
use tokio::sync::mpsc;
|
||||
use tokio_stream::StreamExt;
|
||||
|
||||
use crate::{
|
||||
cursor::{
|
||||
protocol::connect::{self, END_STREAM_FLAG},
|
||||
services::observability::CursorTraceRecorder,
|
||||
transport::{TransportHandle, TransportRegistry},
|
||||
},
|
||||
Result,
|
||||
};
|
||||
|
||||
pub async fn stream(registry: &TransportRegistry, request_id: &str) -> Result<Response<Body>> {
|
||||
let handle = registry.get_or_create(request_id).await?;
|
||||
let receiver = handle.subscribe();
|
||||
let trace = handle.trace().cloned();
|
||||
if let Some(trace) = &trace {
|
||||
trace.response_started(StatusCode::OK.as_u16()).await;
|
||||
}
|
||||
let body_stream = local_body_stream(receiver, handle, trace);
|
||||
let mut response = Response::new(Body::from_stream(body_stream));
|
||||
*response.status_mut() = StatusCode::OK;
|
||||
response.headers_mut().insert(
|
||||
header::CONTENT_TYPE,
|
||||
HeaderValue::from_static("text/event-stream"),
|
||||
);
|
||||
response
|
||||
.headers_mut()
|
||||
.insert(header::CACHE_CONTROL, HeaderValue::from_static("no-cache"));
|
||||
response
|
||||
.headers_mut()
|
||||
.insert("connect-protocol-version", HeaderValue::from_static("1"));
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
fn local_body_stream(
|
||||
mut receiver: mpsc::UnboundedReceiver<Bytes>,
|
||||
handle: TransportHandle,
|
||||
trace: Option<CursorTraceRecorder>,
|
||||
) -> impl tokio_stream::Stream<Item = std::result::Result<Bytes, Infallible>> {
|
||||
async_stream::stream! {
|
||||
let mut guard = LocalRunGuard::new(handle);
|
||||
let mut trace = TraceStreamSink::new(trace, "byok_server");
|
||||
while let Some(chunk) = receiver.recv().await {
|
||||
let terminal = is_end_stream_frame(&chunk);
|
||||
trace.chunk(&chunk);
|
||||
if terminal {
|
||||
guard.complete();
|
||||
trace.finish(end_stream_error(&chunk));
|
||||
}
|
||||
yield Ok::<Bytes, Infallible>(chunk);
|
||||
if terminal {
|
||||
return;
|
||||
}
|
||||
}
|
||||
guard.complete();
|
||||
trace.finish(None);
|
||||
}
|
||||
}
|
||||
|
||||
fn is_end_stream_frame(frame: &Bytes) -> bool {
|
||||
frame
|
||||
.first()
|
||||
.is_some_and(|flags| flags & END_STREAM_FLAG != 0)
|
||||
}
|
||||
|
||||
fn end_stream_error(frame: &Bytes) -> Option<String> {
|
||||
connect::decode_frames(frame)
|
||||
.ok()?
|
||||
.into_iter()
|
||||
.find_map(|(flags, payload)| {
|
||||
if flags & END_STREAM_FLAG == 0 {
|
||||
return None;
|
||||
}
|
||||
let value = serde_json::from_slice::<serde_json::Value>(&payload).ok()?;
|
||||
let error = value.get("error")?;
|
||||
let code = error.get("code").and_then(serde_json::Value::as_str);
|
||||
let message = error
|
||||
.get("message")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.filter(|message| !message.is_empty());
|
||||
Some(match (code, message) {
|
||||
(Some(code), Some(message)) => format!("{code}: {message}"),
|
||||
(Some(code), None) => code.to_string(),
|
||||
(None, Some(message)) => message.to_string(),
|
||||
(None, None) => error.to_string(),
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
struct LocalRunGuard {
|
||||
handle: TransportHandle,
|
||||
completed: bool,
|
||||
}
|
||||
|
||||
impl LocalRunGuard {
|
||||
fn new(handle: TransportHandle) -> Self {
|
||||
Self {
|
||||
handle,
|
||||
completed: false,
|
||||
}
|
||||
}
|
||||
|
||||
fn complete(&mut self) {
|
||||
self.completed = true;
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for LocalRunGuard {
|
||||
fn drop(&mut self) {
|
||||
if !self.completed {
|
||||
let handle = self.handle.clone();
|
||||
tokio::spawn(async move {
|
||||
handle.disconnect().await;
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn upstream(
|
||||
registry: TransportRegistry,
|
||||
request_id: String,
|
||||
generation: u64,
|
||||
response: Response<Body>,
|
||||
trace: Option<CursorTraceRecorder>,
|
||||
) -> Response<Body> {
|
||||
let (parts, body) = response.into_parts();
|
||||
if let Some(trace) = &trace {
|
||||
trace.response_started(parts.status.as_u16()).await;
|
||||
}
|
||||
let stream = async_stream::stream! {
|
||||
let _guard = UpstreamRunGuard {
|
||||
registry,
|
||||
request_id,
|
||||
generation,
|
||||
};
|
||||
let mut trace = TraceStreamSink::new(trace, "cursor_official");
|
||||
let mut body = body.into_data_stream();
|
||||
while let Some(chunk) = body.next().await {
|
||||
match chunk {
|
||||
Ok(chunk) => {
|
||||
trace.chunk(&chunk);
|
||||
yield Ok::<Bytes, axum::Error>(chunk);
|
||||
}
|
||||
Err(error) => {
|
||||
trace.finish(Some(error.to_string()));
|
||||
yield Err(error);
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
trace.finish(None);
|
||||
};
|
||||
Response::from_parts(parts, Body::from_stream(stream))
|
||||
}
|
||||
|
||||
enum TraceStreamEvent {
|
||||
Chunk(Bytes),
|
||||
Finish(Option<String>),
|
||||
}
|
||||
|
||||
struct TraceStreamSink {
|
||||
sender: Option<mpsc::UnboundedSender<TraceStreamEvent>>,
|
||||
}
|
||||
|
||||
impl TraceStreamSink {
|
||||
fn new(trace: Option<CursorTraceRecorder>, source: &'static str) -> Self {
|
||||
let Some(trace) = trace else {
|
||||
return Self { sender: None };
|
||||
};
|
||||
let (sender, mut receiver) = mpsc::unbounded_channel();
|
||||
tokio::spawn(async move {
|
||||
while let Some(event) = receiver.recv().await {
|
||||
match event {
|
||||
TraceStreamEvent::Chunk(chunk) => {
|
||||
trace.response_chunk(source, &chunk).await;
|
||||
}
|
||||
TraceStreamEvent::Finish(error) => {
|
||||
trace.finish(error.as_deref()).await;
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
trace.finish(None).await;
|
||||
});
|
||||
Self {
|
||||
sender: Some(sender),
|
||||
}
|
||||
}
|
||||
|
||||
fn chunk(&self, chunk: &Bytes) {
|
||||
if let Some(sender) = &self.sender {
|
||||
let _ = sender.send(TraceStreamEvent::Chunk(chunk.clone()));
|
||||
}
|
||||
}
|
||||
|
||||
fn finish(&mut self, error: Option<String>) {
|
||||
if let Some(sender) = self.sender.take() {
|
||||
let _ = sender.send(TraceStreamEvent::Finish(error));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for TraceStreamSink {
|
||||
fn drop(&mut self) {
|
||||
if self.sender.is_some() {
|
||||
self.finish(Some(
|
||||
"response stream dropped before completion".to_string(),
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct UpstreamRunGuard {
|
||||
registry: TransportRegistry,
|
||||
request_id: String,
|
||||
generation: u64,
|
||||
}
|
||||
|
||||
impl Drop for UpstreamRunGuard {
|
||||
fn drop(&mut self) {
|
||||
self.registry
|
||||
.finish_upstream(self.request_id.clone(), self.generation);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,6 @@
|
||||
//! Exposes the HTTP and Connect API layer.
|
||||
|
||||
pub mod cursor;
|
||||
mod router;
|
||||
|
||||
pub use router::router;
|
||||
@@ -0,0 +1,7 @@
|
||||
//! Builds the top-level server router.
|
||||
|
||||
use crate::{cursor::transport::TransportRegistry, Result};
|
||||
|
||||
pub fn router(registry: TransportRegistry) -> Result<axum::Router> {
|
||||
super::cursor::router(registry)
|
||||
}
|
||||
Reference in New Issue
Block a user