mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-06 13:44:21 +08:00
237 lines
7.7 KiB
Rust
237 lines
7.7 KiB
Rust
use axum::{
|
|
body::{Body, Bytes},
|
|
extract::Extension,
|
|
http::{header, HeaderValue, Request, Response, StatusCode},
|
|
};
|
|
use base64::{engine::general_purpose::STANDARD, Engine};
|
|
use bytes::{BufMut, BytesMut};
|
|
use prost::Message;
|
|
use serde_json::{json, Map, Value};
|
|
use sha2::{Digest, Sha256};
|
|
|
|
use crate::{cursor::proxy, Error, Result};
|
|
|
|
pub const BOOTSTRAP_STATSIG_PATH: &str = "/aiserver.v1.AnalyticsService/BootstrapStatsig";
|
|
const AGENT_RETRIES_GATE: &str = "nal_agent_retries";
|
|
const LOCAL_RULE: &str = "local_enabled";
|
|
|
|
#[derive(Clone, PartialEq, Message)]
|
|
struct BootstrapStatsigResponse {
|
|
#[prost(string, tag = "1")]
|
|
config: String,
|
|
#[prost(uint64, tag = "2")]
|
|
generated_at_ms: u64,
|
|
}
|
|
|
|
pub async fn bootstrap_statsig(
|
|
Extension(upstream): Extension<proxy::CursorProxy>,
|
|
request: Request<Body>,
|
|
) -> Result<Response<Body>> {
|
|
match proxy::forward_buffered(&upstream, request).await {
|
|
Ok(response) if response.status.is_success() => match patch_upstream(response) {
|
|
Ok(response) => Ok(response),
|
|
Err(error) => {
|
|
tracing::warn!(%error, "Cursor Statsig bootstrap was invalid; using local bootstrap");
|
|
local_response()
|
|
}
|
|
},
|
|
Ok(response) => {
|
|
tracing::warn!(status = %response.status, "Cursor Statsig bootstrap was rejected; using local bootstrap");
|
|
local_response()
|
|
}
|
|
Err(error) => {
|
|
tracing::warn!(%error, "Cursor Statsig bootstrap was unavailable; using local bootstrap");
|
|
local_response()
|
|
}
|
|
}
|
|
}
|
|
|
|
fn patch_upstream(response: proxy::BufferedResponse) -> Result<Response<Body>> {
|
|
let (framed, payload) = unary_payload(&response.body)?;
|
|
let mut message = BootstrapStatsigResponse::decode(payload)?;
|
|
let mut config = serde_json::from_str::<Value>(&message.config)?;
|
|
enable_agent_retries(&mut config)?;
|
|
message.config = serde_json::to_string(&config)?;
|
|
Ok(response.with_body(encode_unary(&message, framed)))
|
|
}
|
|
|
|
fn local_response() -> Result<Response<Body>> {
|
|
let generated_at_ms = chrono::Utc::now().timestamp_millis() as u64;
|
|
let mut config = json!({
|
|
"feature_gates": {},
|
|
"dynamic_configs": {},
|
|
"layer_configs": {},
|
|
"user": {
|
|
"userID": "local_ultra",
|
|
"customIDs": { "localUserID": "local_ultra" }
|
|
},
|
|
"has_updates": true,
|
|
"hash_used": "none",
|
|
"sdkParams": {
|
|
"stableID": "local_ultra",
|
|
"disableDiagnosticsLogging": true
|
|
},
|
|
"time": generated_at_ms
|
|
});
|
|
enable_agent_retries(&mut config)?;
|
|
let message = BootstrapStatsigResponse {
|
|
config: serde_json::to_string(&config)?,
|
|
generated_at_ms,
|
|
};
|
|
let body = message.encode_to_vec();
|
|
let mut response = Response::new(Body::from(body.clone()));
|
|
*response.status_mut() = StatusCode::OK;
|
|
response.headers_mut().insert(
|
|
header::CONTENT_TYPE,
|
|
HeaderValue::from_static("application/proto"),
|
|
);
|
|
response.headers_mut().insert(
|
|
header::CONTENT_LENGTH,
|
|
body.len()
|
|
.to_string()
|
|
.parse()
|
|
.expect("body length is a valid header value"),
|
|
);
|
|
Ok(response)
|
|
}
|
|
|
|
fn enable_agent_retries(config: &mut Value) -> Result<()> {
|
|
let gate_key = statsig_key(config, AGENT_RETRIES_GATE);
|
|
let root = config
|
|
.as_object_mut()
|
|
.ok_or_else(|| Error::Protocol("Statsig bootstrap config must be an object".into()))?;
|
|
let gates = root
|
|
.entry("feature_gates")
|
|
.or_insert_with(|| Value::Object(Map::new()))
|
|
.as_object_mut()
|
|
.ok_or_else(|| Error::Protocol("Statsig feature_gates must be an object".into()))?;
|
|
gates.insert(gate_key.clone(), enabled_gate(&gate_key));
|
|
Ok(())
|
|
}
|
|
|
|
fn statsig_key(config: &Value, name: &str) -> String {
|
|
match config.get("hash_used").and_then(Value::as_str) {
|
|
Some("djb2") => djb2(name),
|
|
Some("sha256") => STANDARD.encode(Sha256::digest(name.as_bytes())),
|
|
_ => name.to_owned(),
|
|
}
|
|
}
|
|
|
|
fn djb2(value: &str) -> String {
|
|
value
|
|
.encode_utf16()
|
|
.fold(0_u32, |hash, character| {
|
|
hash.wrapping_mul(31).wrapping_add(u32::from(character))
|
|
})
|
|
.to_string()
|
|
}
|
|
|
|
fn enabled_gate(name: &str) -> Value {
|
|
json!({
|
|
"name": name,
|
|
"value": true,
|
|
"rule_id": LOCAL_RULE,
|
|
"ruleID": LOCAL_RULE,
|
|
"group_name": LOCAL_RULE,
|
|
"groupName": LOCAL_RULE,
|
|
"secondary_exposures": [],
|
|
"secondaryExposures": [],
|
|
"undelegated_secondary_exposures": [],
|
|
"undelegatedSecondaryExposures": [],
|
|
"is_device_based": false,
|
|
"isDeviceBased": false,
|
|
"id_type": "userID",
|
|
"idType": "userID"
|
|
})
|
|
}
|
|
|
|
fn unary_payload(body: &Bytes) -> Result<(bool, &[u8])> {
|
|
if body.len() < 5 {
|
|
return Ok((false, body));
|
|
}
|
|
let flags = body[0];
|
|
let length = u32::from_be_bytes([body[1], body[2], body[3], body[4]]) as usize;
|
|
if length != body.len() - 5 {
|
|
return Ok((false, body));
|
|
}
|
|
if flags != 0 {
|
|
return Err(Error::Protocol(format!(
|
|
"cannot patch compressed or terminal Statsig frame: flags={flags}"
|
|
)));
|
|
}
|
|
Ok((true, &body[5..]))
|
|
}
|
|
|
|
fn encode_unary(message: &impl Message, framed: bool) -> Bytes {
|
|
let payload = message.encode_to_vec();
|
|
if !framed {
|
|
return Bytes::from(payload);
|
|
}
|
|
let mut output = BytesMut::with_capacity(5 + payload.len());
|
|
output.put_u8(0);
|
|
output.put_u32(payload.len() as u32);
|
|
output.extend_from_slice(&payload);
|
|
output.freeze()
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
|
|
#[test]
|
|
fn overlays_retry_gate_without_losing_upstream_config() {
|
|
let mut config = json!({
|
|
"feature_gates": {
|
|
"upstream_gate": { "name": "upstream_gate", "value": true }
|
|
},
|
|
"dynamic_configs": { "kept": { "value": 1 } }
|
|
});
|
|
|
|
enable_agent_retries(&mut config).unwrap();
|
|
|
|
assert_eq!(config["feature_gates"][AGENT_RETRIES_GATE]["value"], true);
|
|
assert_eq!(config["feature_gates"]["upstream_gate"]["value"], true);
|
|
assert_eq!(config["dynamic_configs"]["kept"]["value"], 1);
|
|
}
|
|
|
|
#[test]
|
|
fn uses_the_hash_algorithm_declared_by_upstream() {
|
|
let mut config = json!({
|
|
"hash_used": "djb2",
|
|
"feature_gates": {}
|
|
});
|
|
|
|
enable_agent_retries(&mut config).unwrap();
|
|
|
|
let key = djb2(AGENT_RETRIES_GATE);
|
|
assert_eq!(config["feature_gates"][&key]["name"], key);
|
|
assert_eq!(config["feature_gates"][&key]["value"], true);
|
|
assert!(config["feature_gates"].get(AGENT_RETRIES_GATE).is_none());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn patches_raw_and_connect_framed_responses() {
|
|
for framed in [false, true] {
|
|
let message = BootstrapStatsigResponse {
|
|
config: json!({ "feature_gates": {} }).to_string(),
|
|
generated_at_ms: 123,
|
|
};
|
|
let body = encode_unary(&message, framed);
|
|
let buffered = proxy::BufferedResponse {
|
|
status: StatusCode::OK,
|
|
headers: Default::default(),
|
|
body,
|
|
};
|
|
|
|
let response = patch_upstream(buffered).unwrap();
|
|
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
|
|
.await
|
|
.unwrap();
|
|
let (_, payload) = unary_payload(&body).unwrap();
|
|
let patched = BootstrapStatsigResponse::decode(payload).unwrap();
|
|
let config: Value = serde_json::from_str(&patched.config).unwrap();
|
|
assert_eq!(config["feature_gates"][AGENT_RETRIES_GATE]["value"], true);
|
|
}
|
|
}
|
|
}
|