use crate::search::SearchMetrics;
use serde::{Deserialize, Serialize};
pub const BINARY_MAGIC: u8 = 0x00;
pub const BINARY_PROTOCOL_VERSION: u8 = 3;
#[derive(Debug)]
pub enum BinaryFrameError {
UnsupportedVersion { expected: u8, got: u8 },
Decode(postcard::Error),
}
impl std::fmt::Display for BinaryFrameError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::UnsupportedVersion { expected, got } => {
write!(
f,
"unsupported binary protocol version: expected {expected}, got {got}"
)
}
Self::Decode(e) => write!(f, "postcard decode error: {e}"),
}
}
}
impl std::error::Error for BinaryFrameError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::UnsupportedVersion { .. } => None,
Self::Decode(e) => Some(e),
}
}
}
impl From<postcard::Error> for BinaryFrameError {
fn from(e: postcard::Error) -> Self {
Self::Decode(e)
}
}
#[derive(Debug, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum Request {
Search(SearchRequest),
Health,
Shutdown,
GraphWalk(GraphWalkRequest),
MultiSearch(MultiSearchRequest),
DeepSearch(DeepSearchRequest),
Agent(AgentRequest),
AgentHits(AgentRequest),
}
#[derive(Debug, Serialize, Deserialize, Clone)]
pub struct SearchRequest {
pub query: String,
#[serde(default = "default_max_results")]
pub max_results: usize,
#[serde(default = "default_true")]
pub use_dense: bool,
#[serde(default = "default_true")]
pub use_sparse: bool,
#[serde(default)]
pub use_rerank: bool,
#[serde(default)]
pub include_types: Vec<String>,
#[serde(default)]
pub exclude_types: Vec<String>,
#[serde(default)]
pub code_only: bool,
#[serde(default = "default_true")]
pub include_content: bool,
#[serde(default)]
pub snippet: bool,
#[serde(default)]
pub grep_mode: bool,
#[serde(default)]
pub regex_pattern: Option<String>,
#[serde(default)]
pub auto_peek_top: bool,
}
fn default_max_results() -> usize {
10
}
fn default_true() -> bool {
true
}
fn default_deep_max_results() -> usize {
20
}
#[derive(Debug, Serialize, Deserialize, Clone)]
pub struct DeepSearchRequest {
pub query: String,
#[serde(default = "default_deep_max_results")]
pub max_results: usize,
#[serde(default = "default_true")]
pub use_graph: bool,
}
#[derive(Debug, Serialize, Deserialize, Clone)]
pub struct DeepSearchSource {
pub file: String,
pub start_line: u32,
pub end_line: u32,
#[serde(default)]
pub name: Option<String>,
#[serde(default)]
pub kind: Option<String>,
}
#[derive(Debug, Serialize, Deserialize, Clone, Default)]
pub struct DeepResponseMetrics {
pub search_ms: u64,
pub triage_ms: u64,
pub graph_ms: u64,
pub read_ms: u64,
pub summarize_ms: u64,
pub total_ms: u64,
pub chunks_searched: usize,
pub chunks_read: usize,
#[serde(default)]
pub confidence_zone: String,
}
#[derive(Debug, Serialize, Deserialize, Clone)]
pub struct DeepSearchResponse {
pub answer: String,
pub sources: Vec<DeepSearchSource>,
pub metrics: DeepResponseMetrics,
#[serde(default)]
pub confidence: f32,
}
#[derive(Debug, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum Response {
Search(SearchResponse),
Health(HealthResponse),
Shutdown(ShutdownResponse),
Error(ErrorResponse),
GraphWalk(GraphWalkResponse),
MultiSearch(MultiSearchResponse),
DeepSearch(DeepSearchResponse),
Agent(AgentResponse),
AgentHits(AgentHitsResponse),
}
#[derive(Debug, Serialize, Deserialize, Clone)]
pub struct SearchResponse {
pub results: Vec<SearchResultItem>,
pub duration_ms: u64,
pub dense_count: usize,
pub sparse_count: usize,
pub fused_count: usize,
#[serde(default)]
pub metrics: Option<SearchMetrics>,
#[serde(default)]
pub confidence: Option<String>,
#[serde(default)]
pub disambiguation: Option<Vec<DisambigSuggestion>>,
}
#[derive(Debug, Serialize, Deserialize, Clone, PartialEq)]
pub struct DisambigSuggestion {
pub name: String,
pub path: String,
pub line: usize,
}
#[derive(Debug, Serialize, Deserialize, Clone)]
pub struct SearchResultItem {
pub file: String,
pub start_line: u32,
pub end_line: u32,
pub score: f32,
pub source: String,
pub chunk_type: String,
#[serde(default)]
pub name: Option<String>,
#[serde(default)]
pub language: Option<String>,
#[serde(default)]
pub content: Option<String>,
#[serde(default)]
pub kind: Option<String>,
#[serde(default)]
pub summary: Option<String>,
}
#[derive(Debug, Serialize, Deserialize)]
pub struct HealthResponse {
pub status: String,
pub uptime_s: u64,
pub searches: u64,
}
#[derive(Debug, Serialize, Deserialize)]
pub struct ShutdownResponse {
pub status: String,
}
#[derive(Debug, Serialize, Deserialize)]
pub struct ErrorResponse {
pub message: String,
}
#[derive(Debug, Serialize, Deserialize, Clone)]
pub struct MultiSearchRequest {
pub queries: Vec<SearchRequest>,
}
#[derive(Debug, Serialize, Deserialize, Clone)]
pub struct MultiSearchResponse {
pub responses: Vec<SearchResponse>,
}
#[derive(Debug, Serialize, Deserialize, Clone)]
pub struct GraphWalkRequest {
pub symbol: String,
}
#[derive(Debug, Serialize, Deserialize, Clone)]
pub struct GraphWalkResponse {
pub target: Vec<SearchResultItem>,
pub callers: Vec<SearchResultItem>,
pub callees: Vec<SearchResultItem>,
pub type_refs: Vec<SearchResultItem>,
pub hierarchy: Vec<SearchResultItem>,
}
#[derive(Debug, Serialize, Deserialize)]
pub enum BinaryRequest {
Search(SearchRequest),
Health,
Shutdown,
GraphWalk(GraphWalkRequest),
MultiSearch(MultiSearchRequest),
DeepSearch(DeepSearchRequest),
Agent(AgentRequest),
AgentHits(AgentRequest),
}
#[derive(Debug, Serialize, Deserialize)]
pub enum BinaryResponse {
Search(SearchResponse),
Health(HealthResponse),
Shutdown(ShutdownResponse),
Error(ErrorResponse),
GraphWalk(GraphWalkResponse),
MultiSearch(MultiSearchResponse),
DeepSearch(DeepSearchResponse),
Agent(AgentResponse),
AgentHits(AgentHitsResponse),
}
impl From<BinaryRequest> for Request {
fn from(br: BinaryRequest) -> Self {
match br {
BinaryRequest::Search(s) => Request::Search(s),
BinaryRequest::Health => Request::Health,
BinaryRequest::Shutdown => Request::Shutdown,
BinaryRequest::GraphWalk(g) => Request::GraphWalk(g),
BinaryRequest::MultiSearch(m) => Request::MultiSearch(m),
BinaryRequest::DeepSearch(d) => Request::DeepSearch(d),
BinaryRequest::Agent(a) => Request::Agent(a),
BinaryRequest::AgentHits(a) => Request::AgentHits(a),
}
}
}
impl From<Response> for BinaryResponse {
fn from(r: Response) -> Self {
match r {
Response::Search(s) => BinaryResponse::Search(s),
Response::Health(h) => BinaryResponse::Health(h),
Response::Shutdown(s) => BinaryResponse::Shutdown(s),
Response::Error(e) => BinaryResponse::Error(e),
Response::GraphWalk(g) => BinaryResponse::GraphWalk(g),
Response::MultiSearch(m) => BinaryResponse::MultiSearch(m),
Response::DeepSearch(d) => BinaryResponse::DeepSearch(d),
Response::Agent(a) => BinaryResponse::Agent(a),
Response::AgentHits(a) => BinaryResponse::AgentHits(a),
}
}
}
pub fn encode_binary_request(req: &BinaryRequest) -> Vec<u8> {
let payload = postcard::to_stdvec(req).expect("postcard serialize");
let body_len = 1 + payload.len();
let len = body_len as u32;
let mut buf = Vec::with_capacity(BINARY_FRAME_HEADER_LEN + payload.len());
buf.push(BINARY_MAGIC);
buf.extend_from_slice(&len.to_le_bytes());
buf.push(BINARY_PROTOCOL_VERSION);
buf.extend_from_slice(&payload);
buf
}
pub fn encode_binary_response(resp: &BinaryResponse) -> Vec<u8> {
let payload = postcard::to_stdvec(resp).expect("postcard serialize");
let body_len = 1 + payload.len();
let len = body_len as u32;
let mut buf = Vec::with_capacity(BINARY_FRAME_HEADER_LEN + payload.len());
buf.push(BINARY_MAGIC);
buf.extend_from_slice(&len.to_le_bytes());
buf.push(BINARY_PROTOCOL_VERSION);
buf.extend_from_slice(&payload);
buf
}
pub const BINARY_FRAME_HEADER_LEN: usize = 1 + 4;
pub fn decode_binary_response(data: &[u8]) -> Result<BinaryResponse, BinaryFrameError> {
let Some((&version, payload)) = data.split_first() else {
return Err(BinaryFrameError::UnsupportedVersion {
expected: BINARY_PROTOCOL_VERSION,
got: 0,
});
};
if version != BINARY_PROTOCOL_VERSION {
return Err(BinaryFrameError::UnsupportedVersion {
expected: BINARY_PROTOCOL_VERSION,
got: version,
});
}
Ok(postcard::from_bytes::<BinaryResponse>(payload)?)
}
pub fn decode_binary_request(data: &[u8]) -> Result<BinaryRequest, BinaryFrameError> {
let Some((&version, payload)) = data.split_first() else {
return Err(BinaryFrameError::UnsupportedVersion {
expected: BINARY_PROTOCOL_VERSION,
got: 0,
});
};
if version != BINARY_PROTOCOL_VERSION {
return Err(BinaryFrameError::UnsupportedVersion {
expected: BINARY_PROTOCOL_VERSION,
got: version,
});
}
Ok(postcard::from_bytes::<BinaryRequest>(payload)?)
}
use crate::search::agent_classifier::AgentRoute;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AgentRequest {
pub query: String,
#[serde(default)]
pub route: Option<AgentRoute>,
#[serde(default)]
pub budget: Option<usize>,
#[serde(default)]
pub full_code: bool,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AgentResponse {
pub route: AgentRoute,
pub formatted: String,
pub metrics: AgentMetrics,
#[serde(default)]
pub disambiguation: Option<Vec<DisambigSuggestion>>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AgentHitsResponse {
pub route: AgentRoute,
pub hits: Vec<SearchResultItem>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AgentMetrics {
pub classify_us: u64,
pub search_ms: u64,
pub format_ms: u64,
pub total_ms: u64,
pub fallback_used: bool,
pub result_count: usize,
}
#[cfg(test)]
mod agent_protocol_tests {
use super::*;
#[test]
fn test_agent_request_json_roundtrip() {
let req = AgentRequest {
query: "how does auth work?".into(),
route: None,
budget: Some(8000),
full_code: false,
};
let json = serde_json::to_string(&req).unwrap();
let parsed: AgentRequest = serde_json::from_str(&json).unwrap();
assert_eq!(parsed.query, "how does auth work?");
assert_eq!(parsed.budget, Some(8000));
}
#[test]
fn test_agent_request_with_route() {
use crate::search::agent_classifier::AgentRoute;
let json = r#"{"query":"AuthService","route":"exact_symbol"}"#;
let req: AgentRequest = serde_json::from_str(json).unwrap();
assert_eq!(req.route, Some(AgentRoute::ExactSymbol));
}
#[test]
fn test_agent_response_json_roundtrip() {
use crate::search::agent_classifier::AgentRoute;
let resp = AgentResponse {
route: AgentRoute::Semantic,
formatted: "[route: semantic]\n\nresults here".into(),
metrics: AgentMetrics {
classify_us: 5,
search_ms: 17,
format_ms: 0,
total_ms: 18,
fallback_used: false,
result_count: 3,
},
disambiguation: None,
};
let json = serde_json::to_string(&resp).unwrap();
let parsed: AgentResponse = serde_json::from_str(&json).unwrap();
assert_eq!(parsed.route, AgentRoute::Semantic);
assert!(parsed.formatted.contains("[route: semantic]"));
}
#[test]
fn test_agent_hits_binary_roundtrip() {
use crate::search::agent_classifier::AgentRoute;
let req = BinaryRequest::AgentHits(AgentRequest {
query: "auth".into(),
route: Some(AgentRoute::Semantic),
budget: None,
full_code: false,
});
let enc = encode_binary_request(&req);
let dec = decode_binary_request(&enc[BINARY_FRAME_HEADER_LEN..]).unwrap();
match dec {
BinaryRequest::AgentHits(r) => {
assert_eq!(r.query, "auth");
assert_eq!(r.route, Some(AgentRoute::Semantic));
}
_ => panic!("Wrong request variant"),
}
let resp = BinaryResponse::AgentHits(AgentHitsResponse {
route: AgentRoute::Semantic,
hits: vec![SearchResultItem {
file: "auth/login.rs".into(),
start_line: 1,
end_line: 3,
score: 0.9,
source: "Hybrid".into(),
chunk_type: "AstNode".into(),
name: Some("authenticate".into()),
language: Some("rust".into()),
content: None,
kind: Some("function".into()),
summary: None,
}],
});
let enc = encode_binary_response(&resp);
let dec = decode_binary_response(&enc[BINARY_FRAME_HEADER_LEN..]).unwrap();
match dec {
BinaryResponse::AgentHits(r) => {
assert_eq!(r.route, AgentRoute::Semantic);
assert_eq!(r.hits.len(), 1);
assert_eq!(r.hits[0].file, "auth/login.rs");
}
_ => panic!("Wrong response variant"),
}
}
#[test]
fn appending_agent_hits_keeps_existing_discriminants() {
let enc = encode_binary_request(&BinaryRequest::Health);
let body = &enc[BINARY_FRAME_HEADER_LEN..];
assert_eq!(body[0], BINARY_PROTOCOL_VERSION);
assert_eq!(body[1], 1, "Health discriminant must stay at index 1");
let agent = encode_binary_request(&BinaryRequest::Agent(AgentRequest {
query: String::new(),
route: None,
budget: None,
full_code: false,
}));
assert_eq!(
agent[BINARY_FRAME_HEADER_LEN + 1],
6,
"Agent discriminant must stay at index 6"
);
}
#[test]
fn test_agent_binary_roundtrip() {
let req = BinaryRequest::Agent(AgentRequest {
query: "test query".into(),
route: None,
budget: None,
full_code: false,
});
let encoded = encode_binary_request(&req);
assert_eq!(encoded[0], BINARY_MAGIC);
assert_eq!(encoded[BINARY_FRAME_HEADER_LEN], BINARY_PROTOCOL_VERSION);
let decoded = decode_binary_request(&encoded[BINARY_FRAME_HEADER_LEN..]).unwrap();
match decoded {
BinaryRequest::Agent(r) => assert_eq!(r.query, "test query"),
_ => panic!("Wrong variant"),
}
}
#[test]
fn unsupported_version_rejected_on_decode() {
let bogus_version = 99u8;
let postcard_payload = postcard::to_stdvec(&BinaryRequest::Health).unwrap();
let mut body = vec![bogus_version];
body.extend_from_slice(&postcard_payload);
let err = decode_binary_request(&body).expect_err("must reject");
match err {
BinaryFrameError::UnsupportedVersion { expected, got } => {
assert_eq!(expected, BINARY_PROTOCOL_VERSION);
assert_eq!(got, bogus_version);
}
BinaryFrameError::Decode(e) => panic!("expected version mismatch, got decode: {e}"),
}
}
#[test]
fn unsupported_version_rejected_on_response_decode() {
let postcard_payload = postcard::to_stdvec(&BinaryResponse::Health(HealthResponse {
status: "ok".into(),
uptime_s: 0,
searches: 0,
}))
.unwrap();
let mut body = vec![42u8];
body.extend_from_slice(&postcard_payload);
let err = decode_binary_response(&body).expect_err("must reject");
assert!(matches!(
err,
BinaryFrameError::UnsupportedVersion {
expected: BINARY_PROTOCOL_VERSION,
got: 42
}
));
}
#[test]
fn empty_frame_rejected_cleanly() {
let err = decode_binary_request(&[]).expect_err("must reject");
assert!(matches!(
err,
BinaryFrameError::UnsupportedVersion {
expected: BINARY_PROTOCOL_VERSION,
got: 0
}
));
}
#[test]
fn decode_v1_response_as_v3_rejected_cleanly() {
let postcard_payload = postcard::to_stdvec(&BinaryResponse::Health(HealthResponse {
status: "ok".into(),
uptime_s: 0,
searches: 0,
}))
.unwrap();
let mut v1_body = vec![1u8]; v1_body.extend_from_slice(&postcard_payload);
let err = decode_binary_response(&v1_body).expect_err("must reject v1 frame under v3");
match err {
BinaryFrameError::UnsupportedVersion { expected, got } => {
assert_eq!(expected, BINARY_PROTOCOL_VERSION);
assert_eq!(
expected, 3,
"route-set cutover raised protocol version to 3"
);
assert_eq!(got, 1, "v1 client/daemon payload should be flagged as v1");
}
BinaryFrameError::Decode(e) => {
panic!("v1 frame must surface as UnsupportedVersion, not silent decode error: {e}")
}
}
}
#[test]
fn decode_v2_request_as_v3_rejected_cleanly() {
let postcard_payload = postcard::to_stdvec(&BinaryRequest::Health).unwrap();
let mut v2_body = vec![2u8]; v2_body.extend_from_slice(&postcard_payload);
let err = decode_binary_request(&v2_body).expect_err("must reject v2 frame under v3");
assert!(matches!(
err,
BinaryFrameError::UnsupportedVersion {
expected: 3,
got: 2
}
));
}
#[test]
fn protocol_version_is_v3_for_route_set_cutover() {
assert_eq!(BINARY_PROTOCOL_VERSION, 3);
}
#[test]
fn disambig_suggestion_roundtrips() {
let s = DisambigSuggestion {
name: "userAuthHandler".into(),
path: "auth/users.rs".into(),
line: 42,
};
let json = serde_json::to_string(&s).unwrap();
let parsed: DisambigSuggestion = serde_json::from_str(&json).unwrap();
assert_eq!(parsed, s);
let bytes = postcard::to_stdvec(&s).unwrap();
let decoded: DisambigSuggestion = postcard::from_bytes(&bytes).unwrap();
assert_eq!(decoded, s);
}
}