mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-04 19:31:28 +08:00
feat: add configurable Cursor TAB routing
This commit is contained in:
@@ -155,6 +155,10 @@ pub fn api_router(service: ControlService) -> Router {
|
||||
"/__byok-api__/api/settings/proxy",
|
||||
get(settings::get_proxy).put(settings::update_proxy),
|
||||
)
|
||||
.route(
|
||||
"/__byok-api__/api/settings/tab",
|
||||
get(settings::get_tab).put(settings::update_tab),
|
||||
)
|
||||
.route(
|
||||
"/__byok-api__/api/harness/cursor/status",
|
||||
get(harness::status),
|
||||
|
||||
@@ -17,7 +17,9 @@ use crate::{
|
||||
LlmCallSummary, Overview, ProviderEndpoint, ProviderEndpointInput, ProviderEndpointSecret,
|
||||
ProviderModel, ProviderModelInput, ProviderType,
|
||||
},
|
||||
store::{PortSettings, ProxySettings, ProxySettingsInput, StatisticsStorage, Store},
|
||||
store::{
|
||||
PortSettings, ProxySettings, ProxySettingsInput, StatisticsStorage, Store, TabSettings,
|
||||
},
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
@@ -372,6 +374,14 @@ impl ControlService {
|
||||
pub async fn set_proxy_settings(&self, settings: ProxySettingsInput) -> Result<ProxySettings> {
|
||||
self.store.set_proxy_settings(settings).await
|
||||
}
|
||||
|
||||
pub async fn tab_settings(&self) -> Result<TabSettings> {
|
||||
self.store.tab_settings().await
|
||||
}
|
||||
|
||||
pub async fn set_tab_settings(&self, settings: TabSettings) -> Result<TabSettings> {
|
||||
self.cursor_harness.set_tab_settings(settings).await
|
||||
}
|
||||
}
|
||||
|
||||
fn official_call(trace: CursorRunTraceSummary) -> CallSummary {
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
use crate::Result;
|
||||
use axum::{extract::State, Json};
|
||||
|
||||
use crate::store::{PortSettings, ProxySettings, ProxySettingsInput, StatisticsStorage};
|
||||
use crate::store::{
|
||||
PortSettings, ProxySettings, ProxySettingsInput, StatisticsStorage, TabSettings,
|
||||
};
|
||||
|
||||
use super::{ControlService, ObservabilitySettings};
|
||||
|
||||
@@ -47,3 +49,14 @@ pub async fn update_proxy(
|
||||
) -> Result<Json<ProxySettings>> {
|
||||
Ok(Json(service.set_proxy_settings(settings).await?))
|
||||
}
|
||||
|
||||
pub async fn get_tab(State(service): State<ControlService>) -> Result<Json<TabSettings>> {
|
||||
Ok(Json(service.tab_settings().await?))
|
||||
}
|
||||
|
||||
pub async fn update_tab(
|
||||
State(service): State<ControlService>,
|
||||
Json(settings): Json<TabSettings>,
|
||||
) -> Result<Json<TabSettings>> {
|
||||
Ok(Json(service.set_tab_settings(settings).await?))
|
||||
}
|
||||
|
||||
@@ -13,7 +13,7 @@ use crate::{
|
||||
observability::CursorTraceRecorder,
|
||||
proto::{agent::v1 as agent, aiserver::v1 as ai},
|
||||
proxy::{self, CursorProxy},
|
||||
run_sse,
|
||||
run_sse, tab,
|
||||
},
|
||||
cursor::{CursorParent, CursorSessionRegistry},
|
||||
Result,
|
||||
@@ -70,6 +70,7 @@ fn router_with_proxy(registry: CursorSessionRegistry, proxy: CursorProxy) -> Rou
|
||||
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)
|
||||
|
||||
@@ -22,6 +22,7 @@ pub mod request;
|
||||
pub mod run_sse;
|
||||
pub mod session;
|
||||
pub mod sessions;
|
||||
pub(crate) mod tab;
|
||||
pub mod tools;
|
||||
mod usage;
|
||||
|
||||
|
||||
@@ -82,14 +82,34 @@ impl CursorProxy {
|
||||
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());
|
||||
let url = upstream_url(&parts.headers, &proxy.upstream, path)?;
|
||||
.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);
|
||||
@@ -239,7 +259,7 @@ mod tests {
|
||||
};
|
||||
use tower::ServiceExt;
|
||||
|
||||
use super::{forward, CursorProxy};
|
||||
use super::{forward, forward_to_service, CursorProxy};
|
||||
|
||||
#[tokio::test]
|
||||
async fn preserves_request_and_response() {
|
||||
@@ -289,4 +309,34 @@ mod tests {
|
||||
);
|
||||
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();
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,67 @@
|
||||
use axum::{
|
||||
body::Body,
|
||||
extract::{Extension, State},
|
||||
http::{Request, Response},
|
||||
routing::post,
|
||||
Router,
|
||||
};
|
||||
|
||||
use crate::{
|
||||
cursor::{proxy, CursorSessionRegistry},
|
||||
Result,
|
||||
};
|
||||
|
||||
pub const TAB_PATHS: [&str; 17] = [
|
||||
"/aiserver.v1.AiService/StreamCpp",
|
||||
"/aiserver.v1.AiService/StreamNextCursorPrediction",
|
||||
"/aiserver.v1.AiService/GetCppEditClassification",
|
||||
"/aiserver.v1.AiService/RefreshTabContext",
|
||||
"/aiserver.v1.AiService/CppConfig",
|
||||
"/aiserver.v1.AiService/CppEditHistoryStatus",
|
||||
"/aiserver.v1.AiService/CppAppend",
|
||||
"/aiserver.v1.AiService/CppEditHistoryAppend",
|
||||
"/aiserver.v1.AiService/ReportAiCodeChangeMetrics",
|
||||
"/aiserver.v1.AiService/WriteGitCommitMessage",
|
||||
"/aiserver.v1.AiService/WriteGitBranchName",
|
||||
"/aiserver.v1.CppService/AvailableModels",
|
||||
"/aiserver.v1.CppService/RecordCppFate",
|
||||
"/aiserver.v1.FileSyncService/FSSyncFile",
|
||||
"/aiserver.v1.FileSyncService/FSIsEnabledForUser",
|
||||
"/aiserver.v1.FileSyncService/FSConfig",
|
||||
"/aiserver.v1.FileSyncService/FSUploadFile",
|
||||
];
|
||||
|
||||
pub fn is_tab_path(path: &str) -> bool {
|
||||
TAB_PATHS.contains(&path)
|
||||
}
|
||||
|
||||
pub fn router() -> Router<CursorSessionRegistry> {
|
||||
TAB_PATHS.into_iter().fold(Router::new(), |router, path| {
|
||||
router.route(path, post(forward))
|
||||
})
|
||||
}
|
||||
|
||||
async fn forward(
|
||||
State(registry): State<CursorSessionRegistry>,
|
||||
Extension(upstream): Extension<proxy::CursorProxy>,
|
||||
request: Request<Body>,
|
||||
) -> Result<Response<Body>> {
|
||||
let settings = registry.store().tab_settings().await?;
|
||||
match settings.service_url() {
|
||||
Some(service_url) => proxy::forward_to_service(&upstream, request, service_url).await,
|
||||
None => proxy::forward(Extension(upstream), request).await,
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn matches_only_legacy_tab_routes() {
|
||||
assert_eq!(TAB_PATHS.len(), 17);
|
||||
assert!(is_tab_path("/aiserver.v1.AiService/StreamCpp"));
|
||||
assert!(is_tab_path("/aiserver.v1.FileSyncService/FSUploadFile"));
|
||||
assert!(!is_tab_path("/aiserver.v1.AiService/AvailableModels"));
|
||||
}
|
||||
}
|
||||
@@ -9,7 +9,10 @@ use parking_lot::RwLock;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use tokio::sync::Mutex;
|
||||
|
||||
use crate::{store::Store, Error, Result};
|
||||
use crate::{
|
||||
store::{Store, TabMode, TabSettings},
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
use self::{ca::CaManager, proxy::ProxyRuntime};
|
||||
|
||||
@@ -65,6 +68,7 @@ struct Inner {
|
||||
ca: CaManager,
|
||||
ca_initialization: Mutex<()>,
|
||||
backend_addr: RwLock<Option<SocketAddr>>,
|
||||
tab_mode: Arc<RwLock<TabMode>>,
|
||||
proxy: Mutex<ProxyRuntime>,
|
||||
}
|
||||
|
||||
@@ -76,6 +80,7 @@ impl CursorHarness {
|
||||
ca: CaManager::managed()?,
|
||||
ca_initialization: Mutex::new(()),
|
||||
backend_addr: RwLock::new(None),
|
||||
tab_mode: Arc::new(RwLock::new(TabMode::default())),
|
||||
proxy: Mutex::new(ProxyRuntime::default()),
|
||||
}),
|
||||
})
|
||||
@@ -138,6 +143,12 @@ impl CursorHarness {
|
||||
self.status().await
|
||||
}
|
||||
|
||||
pub async fn set_tab_settings(&self, settings: TabSettings) -> Result<TabSettings> {
|
||||
let saved = self.inner.store.set_tab_settings(settings).await?;
|
||||
*self.inner.tab_mode.write() = saved.mode;
|
||||
Ok(saved)
|
||||
}
|
||||
|
||||
async fn enable(&self) -> Result<()> {
|
||||
if !matches!(self.inner.ca.state()?, CaState::Ready) {
|
||||
return Err(Error::Config(
|
||||
@@ -158,7 +169,15 @@ impl CursorHarness {
|
||||
}
|
||||
let ca = self.inner.ca.load()?;
|
||||
let requested_port = self.inner.store.port_settings().await?.proxy_port;
|
||||
let (url, actual_port) = proxy.start(backend_addr, ca, requested_port).await?;
|
||||
*self.inner.tab_mode.write() = self.inner.store.tab_settings().await?.mode;
|
||||
let (url, actual_port) = proxy
|
||||
.start(
|
||||
backend_addr,
|
||||
ca,
|
||||
requested_port,
|
||||
self.inner.tab_mode.clone(),
|
||||
)
|
||||
.await?;
|
||||
if let Err(error) = self.inner.store.set_proxy_port(actual_port).await {
|
||||
proxy.stop().await;
|
||||
return Err(error);
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
use std::net::SocketAddr;
|
||||
use std::{net::SocketAddr, sync::Arc};
|
||||
|
||||
use hudsucker::{
|
||||
certificate_authority::RcgenAuthority,
|
||||
@@ -8,7 +8,13 @@ use hudsucker::{
|
||||
};
|
||||
use tokio::{net::TcpListener, sync::oneshot, task::JoinHandle};
|
||||
|
||||
use crate::{cursor::proxy::UPSTREAM_URL_HEADER, Error, Result};
|
||||
use parking_lot::RwLock;
|
||||
|
||||
use crate::{
|
||||
cursor::{proxy::UPSTREAM_URL_HEADER, tab::is_tab_path},
|
||||
store::TabMode,
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
use super::ca::LoadedCa;
|
||||
|
||||
@@ -33,6 +39,7 @@ impl ProxyRuntime {
|
||||
backend: SocketAddr,
|
||||
ca: LoadedCa,
|
||||
requested_port: u16,
|
||||
tab_mode: Arc<RwLock<TabMode>>,
|
||||
) -> Result<(String, u16)> {
|
||||
if let Some(url) = self.url() {
|
||||
return Ok((url, self.port.unwrap_or_default()));
|
||||
@@ -45,7 +52,7 @@ impl ProxyRuntime {
|
||||
.with_listener(listener)
|
||||
.with_ca(authority)
|
||||
.with_rustls_connector(aws_lc_rs::default_provider())
|
||||
.with_http_handler(CursorRelay { backend })
|
||||
.with_http_handler(CursorRelay { backend, tab_mode })
|
||||
.with_graceful_shutdown(async move {
|
||||
let _ = done.await;
|
||||
})
|
||||
@@ -89,6 +96,7 @@ async fn bind_proxy_listener(requested_port: u16) -> Result<TcpListener> {
|
||||
#[derive(Clone)]
|
||||
struct CursorRelay {
|
||||
backend: SocketAddr,
|
||||
tab_mode: Arc<RwLock<TabMode>>,
|
||||
}
|
||||
|
||||
impl HttpHandler for CursorRelay {
|
||||
@@ -98,7 +106,8 @@ impl HttpHandler for CursorRelay {
|
||||
mut request: Request<Body>,
|
||||
) -> RequestOrResponse {
|
||||
let original = request.uri().clone();
|
||||
if is_cursor_host(original.host().unwrap_or_default()) && is_local_path(original.path()) {
|
||||
let locally_routed = should_route_locally(original.path(), *self.tab_mode.read());
|
||||
if is_cursor_host(original.host().unwrap_or_default()) && locally_routed {
|
||||
if let Ok(value) = original.to_string().parse() {
|
||||
request.headers_mut().insert(UPSTREAM_URL_HEADER, value);
|
||||
}
|
||||
@@ -157,6 +166,10 @@ fn is_local_path(path: &str) -> bool {
|
||||
)
|
||||
}
|
||||
|
||||
fn should_route_locally(path: &str, tab_mode: TabMode) -> bool {
|
||||
is_local_path(path) || (is_tab_path(path) && tab_mode != TabMode::Direct)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
@@ -179,5 +192,17 @@ mod tests {
|
||||
"/aiserver.v1.AnalyticsService/BootstrapStatsig"
|
||||
));
|
||||
assert!(!is_local_path("/unrelated"));
|
||||
assert!(should_route_locally(
|
||||
"/aiserver.v1.AiService/StreamCpp",
|
||||
TabMode::Public
|
||||
));
|
||||
assert!(should_route_locally(
|
||||
"/aiserver.v1.AiService/StreamCpp",
|
||||
TabMode::Custom
|
||||
));
|
||||
assert!(!should_route_locally(
|
||||
"/aiserver.v1.AiService/StreamCpp",
|
||||
TabMode::Direct
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -6,8 +6,11 @@ use super::{now_ms, Store};
|
||||
|
||||
const PORT_SETTINGS_KEY: &str = "network_ports";
|
||||
const PROXY_SETTINGS_KEY: &str = "outbound_proxy";
|
||||
const TAB_SETTINGS_KEY: &str = "cursor_tab";
|
||||
const INSTALLATION_ID_KEY: &str = "installation_id";
|
||||
|
||||
pub const PUBLIC_TAB_SERVICE_URL: &str = "https://tab.leokun.cn";
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, Deserialize, PartialEq, Eq, Serialize)]
|
||||
pub struct PortSettings {
|
||||
pub proxy_port: u16,
|
||||
@@ -28,6 +31,31 @@ impl ProxyMode {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, Deserialize, PartialEq, Eq, Serialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum TabMode {
|
||||
#[default]
|
||||
Public,
|
||||
Direct,
|
||||
Custom,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, Deserialize, PartialEq, Eq, Serialize)]
|
||||
pub struct TabSettings {
|
||||
pub mode: TabMode,
|
||||
pub address: String,
|
||||
}
|
||||
|
||||
impl TabSettings {
|
||||
pub fn service_url(&self) -> Option<&str> {
|
||||
match self.mode {
|
||||
TabMode::Public => Some(PUBLIC_TAB_SERVICE_URL),
|
||||
TabMode::Direct => None,
|
||||
TabMode::Custom => Some(&self.address),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, Deserialize, PartialEq, Eq, Serialize)]
|
||||
pub struct ProxySettingsInput {
|
||||
pub mode: ProxyMode,
|
||||
@@ -139,6 +167,47 @@ impl Store {
|
||||
self.proxy_settings().await
|
||||
}
|
||||
|
||||
pub async fn tab_settings(&self) -> Result<TabSettings> {
|
||||
let value = sqlx::query_scalar::<_, String>(
|
||||
"SELECT value_json FROM service_settings WHERE setting_key = ?",
|
||||
)
|
||||
.bind(TAB_SETTINGS_KEY)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
value
|
||||
.map(|value| serde_json::from_str(&value).map_err(Into::into))
|
||||
.unwrap_or_else(|| Ok(TabSettings::default()))
|
||||
}
|
||||
|
||||
pub async fn set_tab_settings(&self, mut settings: TabSettings) -> Result<TabSettings> {
|
||||
settings.address = settings.address.trim().trim_end_matches('/').to_owned();
|
||||
if settings.mode == TabMode::Custom {
|
||||
let parsed = url::Url::parse(&settings.address).map_err(|error| {
|
||||
crate::Error::Config(format!("invalid TAB service address: {error}"))
|
||||
})?;
|
||||
if !matches!(parsed.scheme(), "http" | "https") {
|
||||
return Err(crate::Error::Config(
|
||||
"TAB service address must use http or https".into(),
|
||||
));
|
||||
}
|
||||
if parsed.host_str().is_none()
|
||||
|| parsed.query().is_some()
|
||||
|| parsed.fragment().is_some()
|
||||
{
|
||||
return Err(crate::Error::Config(
|
||||
"TAB service address must be a base URL without a query or fragment".into(),
|
||||
));
|
||||
}
|
||||
}
|
||||
sqlx::query("INSERT INTO service_settings(setting_key, value_json, updated_at_ms) VALUES (?, ?, ?) ON CONFLICT(setting_key) DO UPDATE SET value_json = excluded.value_json, updated_at_ms = excluded.updated_at_ms")
|
||||
.bind(TAB_SETTINGS_KEY)
|
||||
.bind(serde_json::to_string(&settings)?)
|
||||
.bind(now_ms())
|
||||
.execute(&self.pool)
|
||||
.await?;
|
||||
Ok(settings)
|
||||
}
|
||||
|
||||
pub async fn port_settings(&self) -> Result<PortSettings> {
|
||||
let value = sqlx::query_scalar::<_, String>(
|
||||
"SELECT value_json FROM service_settings WHERE setting_key = ?",
|
||||
@@ -244,4 +313,28 @@ mod tests {
|
||||
"secret"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn tab_settings_default_to_public_and_validate_custom_urls() {
|
||||
let store = Store::connect("sqlite::memory:").await.unwrap();
|
||||
assert_eq!(store.tab_settings().await.unwrap(), TabSettings::default());
|
||||
|
||||
let saved = store
|
||||
.set_tab_settings(TabSettings {
|
||||
mode: TabMode::Custom,
|
||||
address: " https://tab.example.com/base/ ".into(),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(saved.address, "https://tab.example.com/base");
|
||||
assert_eq!(store.tab_settings().await.unwrap(), saved);
|
||||
|
||||
assert!(store
|
||||
.set_tab_settings(TabSettings {
|
||||
mode: TabMode::Custom,
|
||||
address: "file:///tmp/tab".into(),
|
||||
})
|
||||
.await
|
||||
.is_err());
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user