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

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);
}
}
}