1use std::{
4 collections::{BTreeSet, HashMap, HashSet},
5 time::Instant,
6};
7
8use crate::{
9 api::{
10 ApiError, CodeGraphContextResponse, CodeRepositoryFreshnessDiagnostics,
11 CodeRepositoryFreshnessState, CodeRepositoryQueryResponse, RequestContext,
12 },
13 application::RelayKnowledgeService,
14 domain::{
15 CodeGraphCodeExcerpt, CodeGraphContextBudget, CodeGraphContextPack,
16 CodeGraphContextProvenance, CodeGraphContextRequest, CodeGraphImpactHint, CodeQueryKind,
17 CodeRetrievalHit, CodeRetrievalLayer, CodeRetrievalRequest,
18 },
19};
20
21const ENTRY_QUERY_KINDS: [CodeQueryKind; 3] = [
22 CodeQueryKind::Hybrid,
23 CodeQueryKind::Definition,
24 CodeQueryKind::Symbol,
25];
26const EXPANSION_QUERY_KINDS: [CodeQueryKind; 4] = [
27 CodeQueryKind::References,
28 CodeQueryKind::Callers,
29 CodeQueryKind::Callees,
30 CodeQueryKind::Imports,
31];
32const MAX_CONTEXT_SEEDS: usize = 3;
33const MAX_EXPANSION_LIMIT: usize = 4;
34
35impl RelayKnowledgeService {
36 pub async fn codegraph_context(
38 &self,
39 request: CodeGraphContextRequest,
40 context: RequestContext,
41 ) -> Result<CodeGraphContextResponse, ApiError> {
42 let started = Instant::now();
43 let mut candidate_count = 0usize;
44 let mut entry_points = Vec::new();
45 let mut freshness_parts = Vec::new();
46 let mut context_request = request.clone();
47 let mut primary = None;
48
49 for kind in ENTRY_QUERY_KINDS {
50 let response = self
51 .run_context_query(
52 &context_request,
53 kind,
54 request.query.clone(),
55 request.limit,
56 &context,
57 None,
58 )
59 .await?;
60 candidate_count = candidate_count.saturating_add(response.results.len());
61 push_unique_hits(&mut entry_points, response.results.clone());
62 if kind == CodeQueryKind::Hybrid {
63 context_request = pinned_context_request(&request, &response.scope);
64 primary = Some(response.clone());
65 } else {
66 freshness_parts.push(response.freshness);
67 }
68 }
69
70 let primary = primary.expect("hybrid entry query always runs");
71 let expansion_constraints = context_query_constraints(&context_request)
72 .map_err(|error| ApiError::invalid_argument(error.to_string()))?;
73 let seeds = context_seeds(&entry_points);
74 let mut related_symbols = Vec::new();
75 let mut graph_paths = Vec::new();
76 let mut context_roles = HashMap::new();
77
78 for seed in &seeds {
79 let seed_query = seed_query(seed, &request.query);
80 for kind in EXPANSION_QUERY_KINDS {
81 let response = self
82 .run_context_query(
83 &context_request,
84 kind,
85 seed_query.clone(),
86 request.limit.min(MAX_EXPANSION_LIMIT),
87 &context,
88 Some(&expansion_constraints),
89 )
90 .await?;
91 candidate_count = candidate_count.saturating_add(response.results.len());
92 remember_context_roles(&mut context_roles, kind, &response.results);
93 let results = response.results;
94 if kind == CodeQueryKind::References {
95 push_unique_hits(&mut related_symbols, results);
96 } else {
97 push_unique_hits(&mut graph_paths, results);
98 }
99 freshness_parts.push(response.freshness);
100 }
101 }
102
103 let count_truncated = truncate_hits(&mut entry_points, request.limit)
104 | truncate_hits(&mut related_symbols, request.limit)
105 | truncate_hits(&mut graph_paths, request.limit);
106 apply_code_visibility(&mut entry_points, request.include_code);
107 apply_code_visibility(&mut related_symbols, request.include_code);
108 apply_code_visibility(&mut graph_paths, request.include_code);
109
110 let mut pack = CodeGraphContextPack {
111 code_excerpts: code_excerpts(
112 request.include_code,
113 &entry_points,
114 &related_symbols,
115 &graph_paths,
116 &context_roles,
117 ),
118 impact_hints: impact_hints(&graph_paths, &context_roles),
119 entry_points,
120 related_symbols,
121 graph_paths,
122 };
123 let byte_truncated = pack_to_budget(&mut pack, request.max_context_bytes, &context_roles);
124 let mut truncated = count_truncated | byte_truncated;
125 let context_bytes = serialized_context_bytes(&pack);
126 truncated |= primary.results.len() > request.limit;
127 let retrieval_layers = retrieval_layers(&pack);
128 let mut freshness = merge_context_freshness(primary.freshness, freshness_parts);
129 freshness.merge_direct_source_read_paths(context_paths(&pack));
130 let returned_count = pack.entry_points.len()
131 + pack.related_symbols.len()
132 + pack.graph_paths.len()
133 + pack.code_excerpts.len();
134 let mut diagnostics = vec![format!(
135 "Expanded {} seed(s) through references, callers, callees, and imports with bounded limits.",
136 seeds.len()
137 )];
138 if count_truncated {
139 diagnostics.push("Context pack was truncated to fit the requested limit.".to_owned());
140 }
141 if byte_truncated {
142 diagnostics.push("Context pack was truncated to fit max_context_bytes.".to_owned());
143 }
144
145 Ok(CodeGraphContextResponse {
146 metadata: primary.metadata,
147 query: request.query.clone(),
148 repository_scope: primary.scope,
149 freshness,
150 budget: CodeGraphContextBudget {
151 limit: request.limit,
152 max_context_bytes: request.max_context_bytes,
153 candidate_count,
154 returned_count,
155 context_bytes,
156 elapsed_ms: started.elapsed().as_millis().try_into().unwrap_or(u64::MAX),
157 },
158 truncated,
159 retrieval_layers,
160 request,
161 pack,
162 diagnostics,
163 })
164 }
165
166 async fn run_context_query(
167 &self,
168 request: &CodeGraphContextRequest,
169 kind: CodeQueryKind,
170 query: String,
171 limit: usize,
172 context: &RequestContext,
173 constraints: Option<&CodeRetrievalRequest>,
174 ) -> Result<CodeRepositoryQueryResponse, ApiError> {
175 let mut retrieval = CodeRetrievalRequest::new(
176 query,
177 request.repository.clone(),
178 kind,
179 limit,
180 request.freshness_policy,
181 )
182 .map_err(|error| ApiError::invalid_argument(error.to_string()))?;
183 retrieval.exclude_generated = request.exclude_generated;
184 if let Some(constraints) = constraints {
185 carry_context_filters(&mut retrieval, constraints);
186 }
187
188 self.query_code_repository(retrieval, context.clone()).await
189 }
190}
191
192fn context_query_constraints(
193 request: &CodeGraphContextRequest,
194) -> Result<CodeRetrievalRequest, crate::domain::DomainError> {
195 CodeRetrievalRequest::new(
196 request.query.clone(),
197 request.repository.clone(),
198 CodeQueryKind::Hybrid,
199 request.limit,
200 request.freshness_policy,
201 )
202}
203
204fn carry_context_filters(target: &mut CodeRetrievalRequest, source: &CodeRetrievalRequest) {
205 target.query_language_filters = source.query_language_filters.clone();
206 target.query_path_substrings = source.query_path_substrings.clone();
207}
208
209fn pinned_context_request(
210 original: &CodeGraphContextRequest,
211 primary_scope: &crate::api::CodeRepositoryScopeMetadata,
212) -> CodeGraphContextRequest {
213 let mut pinned = original.clone();
214 if !primary_scope.resolved_commit_sha.is_empty() {
215 pinned.repository.ref_selector = primary_scope.resolved_commit_sha.clone();
216 }
217 pinned
218}
219
220fn context_seeds(entry_points: &[CodeRetrievalHit]) -> Vec<CodeRetrievalHit> {
221 let mut seeds = entry_points.to_vec();
222 seeds.sort_by(|left, right| right.score.total_cmp(&left.score));
223 seeds.truncate(MAX_CONTEXT_SEEDS);
224 seeds
225}
226
227fn push_unique_hits(target: &mut Vec<CodeRetrievalHit>, hits: Vec<CodeRetrievalHit>) {
228 let mut keys = target.iter().map(hit_key).collect::<HashSet<_>>();
229 for hit in hits {
230 if keys.insert(hit_key(&hit)) {
231 target.push(hit);
232 }
233 }
234 target.sort_by(|left, right| right.score.total_cmp(&left.score));
235}
236
237fn truncate_hits(hits: &mut Vec<CodeRetrievalHit>, limit: usize) -> bool {
238 if hits.len() > limit {
239 hits.truncate(limit);
240 true
241 } else {
242 false
243 }
244}
245
246fn remember_context_roles(
247 roles: &mut HashMap<String, CodeQueryKind>,
248 kind: CodeQueryKind,
249 hits: &[CodeRetrievalHit],
250) {
251 for hit in hits {
252 roles.entry(hit_key(hit)).or_insert(kind);
253 }
254}
255
256fn hit_key(hit: &CodeRetrievalHit) -> String {
257 format!(
258 "{}:{}:{}:{}:{}",
259 hit.path,
260 hit.line_range.start,
261 hit.line_range.end,
262 hit.symbol_snapshot_id.as_deref().unwrap_or(""),
263 hit.edge_kind.as_deref().unwrap_or("")
264 )
265}
266
267fn seed_query(hit: &CodeRetrievalHit, fallback: &str) -> String {
268 hit.canonical_symbol_id
269 .as_deref()
270 .and_then(extract_searchable_symbol_tail)
271 .unwrap_or(fallback)
272 .to_owned()
273}
274
275fn extract_searchable_symbol_tail(value: &str) -> Option<&str> {
276 value
277 .rsplit([':', '/', '#', '.', '@', '[', ']'])
278 .find(|part| part.chars().any(|character| character.is_alphanumeric()))
279}
280
281fn apply_code_visibility(hits: &mut [CodeRetrievalHit], include_code: bool) {
282 if include_code {
283 return;
284 }
285 for hit in hits {
286 hit.excerpt.clear();
287 }
288}
289
290fn code_excerpts(
291 include_code: bool,
292 entry_points: &[CodeRetrievalHit],
293 related_symbols: &[CodeRetrievalHit],
294 graph_paths: &[CodeRetrievalHit],
295 context_roles: &HashMap<String, CodeQueryKind>,
296) -> Vec<CodeGraphCodeExcerpt> {
297 if !include_code {
298 return Vec::new();
299 }
300 entry_points
301 .iter()
302 .chain(related_symbols)
303 .chain(graph_paths)
304 .filter(|hit| !hit.excerpt.trim().is_empty())
305 .map(|hit| CodeGraphCodeExcerpt {
306 path: hit.path.clone(),
307 language_id: hit.language_id.clone(),
308 line_range: hit.line_range.clone(),
309 symbol_snapshot_id: hit.symbol_snapshot_id.clone(),
310 provenance: CodeGraphContextProvenance {
311 query_kind: provenance_kind(hit, context_roles),
312 retrieval_layers: hit.retrieval_layers.clone(),
313 score: hit.score,
314 },
315 excerpt: hit.excerpt.clone(),
316 })
317 .collect()
318}
319
320fn impact_hints(
321 graph_paths: &[CodeRetrievalHit],
322 context_roles: &HashMap<String, CodeQueryKind>,
323) -> Vec<CodeGraphImpactHint> {
324 graph_paths
325 .iter()
326 .map(|hit| CodeGraphImpactHint {
327 path: hit.path.clone(),
328 line_range: hit.line_range.clone(),
329 relationship: context_roles
330 .get(&hit_key(hit))
331 .map(context_role_relationship)
332 .or(hit.edge_kind.as_deref())
333 .unwrap_or_else(|| relationship_from_layers(&hit.retrieval_layers))
334 .to_owned(),
335 symbol_snapshot_id: hit.symbol_snapshot_id.clone(),
336 retrieval_layers: hit.retrieval_layers.clone(),
337 score: hit.score,
338 })
339 .collect()
340}
341
342fn provenance_kind(
343 hit: &CodeRetrievalHit,
344 context_roles: &HashMap<String, CodeQueryKind>,
345) -> CodeQueryKind {
346 if let Some(kind) = context_roles.get(&hit_key(hit)) {
347 *kind
348 } else if hit
349 .retrieval_layers
350 .contains(&CodeRetrievalLayer::Reference)
351 {
352 CodeQueryKind::References
353 } else if hit
354 .retrieval_layers
355 .contains(&CodeRetrievalLayer::CallGraph)
356 {
357 CodeQueryKind::Callers
358 } else if hit
359 .retrieval_layers
360 .contains(&CodeRetrievalLayer::ImportGraph)
361 {
362 CodeQueryKind::Imports
363 } else if hit
364 .retrieval_layers
365 .contains(&CodeRetrievalLayer::Definition)
366 {
367 CodeQueryKind::Definition
368 } else if hit.retrieval_layers.contains(&CodeRetrievalLayer::Symbol) {
369 CodeQueryKind::Symbol
370 } else {
371 CodeQueryKind::Hybrid
372 }
373}
374
375fn context_role_relationship(kind: &CodeQueryKind) -> &'static str {
376 match kind {
377 CodeQueryKind::References => "reference",
378 CodeQueryKind::Callers => "caller",
379 CodeQueryKind::Callees => "callee",
380 CodeQueryKind::Imports => "import",
381 _ => "context",
382 }
383}
384
385fn relationship_from_layers(layers: &[CodeRetrievalLayer]) -> &'static str {
386 if layers.contains(&CodeRetrievalLayer::CallGraph) {
387 "call_graph"
388 } else if layers.contains(&CodeRetrievalLayer::ImportGraph) {
389 "import_graph"
390 } else if layers.contains(&CodeRetrievalLayer::Reference) {
391 "reference"
392 } else {
393 "related"
394 }
395}
396
397fn merge_context_freshness(
398 mut primary: CodeRepositoryFreshnessDiagnostics,
399 parts: Vec<CodeRepositoryFreshnessDiagnostics>,
400) -> CodeRepositoryFreshnessDiagnostics {
401 for freshness in parts {
402 primary.state = worse_freshness_state(primary.state, freshness.state);
403 primary.scope_stale |= freshness.scope_stale;
404 primary.direct_source_read_required |= freshness.direct_source_read_required;
405 primary.index_lag.requested_ref_indexed &= freshness.index_lag.requested_ref_indexed;
406 primary.index_lag.pending_task_count = primary
407 .index_lag
408 .pending_task_count
409 .max(freshness.index_lag.pending_task_count);
410 primary.index_lag.pending_file_count = max_optional_usize(
411 primary.index_lag.pending_file_count,
412 freshness.index_lag.pending_file_count,
413 );
414 primary.pending.active_for_repository |= freshness.pending.active_for_repository;
415 primary.pending.active_matches_request |= freshness.pending.active_matches_request;
416 primary.pending.queue_depth = primary
417 .pending
418 .queue_depth
419 .max(freshness.pending.queue_depth);
420 primary.pending.queued_task_count = primary
421 .pending
422 .queued_task_count
423 .max(freshness.pending.queued_task_count);
424 primary.pending.running_task_count = primary
425 .pending
426 .running_task_count
427 .max(freshness.pending.running_task_count);
428 primary.pending.retrying_task_count = primary
429 .pending
430 .retrying_task_count
431 .max(freshness.pending.retrying_task_count);
432 primary.pending.dead_letter_task_count = primary
433 .pending
434 .dead_letter_task_count
435 .max(freshness.pending.dead_letter_task_count);
436 primary.pending.running_lease_count = primary
437 .pending
438 .running_lease_count
439 .max(freshness.pending.running_lease_count);
440 primary.stale_reason = merge_reason(primary.stale_reason.take(), freshness.stale_reason);
441 primary.degraded_reason =
442 merge_reason(primary.degraded_reason.take(), freshness.degraded_reason);
443 primary.merge_direct_source_read_paths(freshness.direct_source_read_paths);
444 }
445 primary
446}
447
448fn worse_freshness_state(
449 left: CodeRepositoryFreshnessState,
450 right: CodeRepositoryFreshnessState,
451) -> CodeRepositoryFreshnessState {
452 if freshness_state_rank(left) >= freshness_state_rank(right) {
453 left
454 } else {
455 right
456 }
457}
458
459fn freshness_state_rank(state: CodeRepositoryFreshnessState) -> u8 {
460 match state {
461 CodeRepositoryFreshnessState::Fresh => 0,
462 CodeRepositoryFreshnessState::Degraded => 1,
463 CodeRepositoryFreshnessState::Stale => 2,
464 CodeRepositoryFreshnessState::Pending => 3,
465 }
466}
467
468fn max_optional_usize(left: Option<usize>, right: Option<usize>) -> Option<usize> {
469 match (left, right) {
470 (Some(left), Some(right)) => Some(left.max(right)),
471 (Some(value), None) | (None, Some(value)) => Some(value),
472 (None, None) => None,
473 }
474}
475
476fn merge_reason(left: Option<String>, right: Option<String>) -> Option<String> {
477 match (left, right) {
478 (Some(left), Some(right)) if left == right => Some(left),
479 (Some(left), Some(right)) => Some(format!("{left}; {right}")),
480 (Some(reason), None) | (None, Some(reason)) => Some(reason),
481 (None, None) => None,
482 }
483}
484
485fn pack_to_budget(
486 pack: &mut CodeGraphContextPack,
487 max_context_bytes: usize,
488 context_roles: &HashMap<String, CodeQueryKind>,
489) -> bool {
490 let mut truncated = false;
491 while serialized_context_bytes(pack) > max_context_bytes {
492 let removed_code_excerpt = pack.code_excerpts.pop().is_some();
493 let cleared_expansion_excerpts =
494 !removed_code_excerpt && clear_expansion_hit_excerpts(pack);
495 if removed_code_excerpt || cleared_expansion_excerpts {
496 truncated = true;
497 } else if pack.graph_paths.pop().is_some() {
498 pack.impact_hints = impact_hints(&pack.graph_paths, context_roles);
499 truncated = true;
500 } else if pack.related_symbols.pop().is_some()
501 || clear_entry_hit_excerpts(pack)
502 || pack.entry_points.pop().is_some()
503 {
504 truncated = true;
505 } else {
506 clear_hit_excerpts(pack);
507 return true;
508 }
509 }
510
511 truncated
512}
513
514fn clear_hit_excerpts(pack: &mut CodeGraphContextPack) {
515 clear_expansion_hit_excerpts(pack);
516 clear_entry_hit_excerpts(pack);
517 pack.code_excerpts.clear();
518 pack.impact_hints.clear();
519}
520
521fn clear_expansion_hit_excerpts(pack: &mut CodeGraphContextPack) -> bool {
522 let had_excerpts = pack
523 .related_symbols
524 .iter()
525 .chain(&pack.graph_paths)
526 .any(|hit| !hit.excerpt.is_empty());
527 apply_code_visibility(&mut pack.related_symbols, false);
528 apply_code_visibility(&mut pack.graph_paths, false);
529
530 had_excerpts
531}
532
533fn clear_entry_hit_excerpts(pack: &mut CodeGraphContextPack) -> bool {
534 let had_excerpts = pack.entry_points.iter().any(|hit| !hit.excerpt.is_empty());
535 apply_code_visibility(&mut pack.entry_points, false);
536
537 had_excerpts
538}
539
540fn retrieval_layers(pack: &CodeGraphContextPack) -> Vec<CodeRetrievalLayer> {
541 let mut layers = BTreeSet::new();
542 for hit in pack
543 .entry_points
544 .iter()
545 .chain(&pack.related_symbols)
546 .chain(&pack.graph_paths)
547 {
548 for layer in &hit.retrieval_layers {
549 layers.insert(layer.as_str());
550 }
551 }
552
553 layers
554 .into_iter()
555 .filter_map(layer_from_str)
556 .collect::<Vec<_>>()
557}
558
559fn layer_from_str(value: &str) -> Option<CodeRetrievalLayer> {
560 match value {
561 "lexical" => Some(CodeRetrievalLayer::Lexical),
562 "symbol" => Some(CodeRetrievalLayer::Symbol),
563 "definition" => Some(CodeRetrievalLayer::Definition),
564 "reference" => Some(CodeRetrievalLayer::Reference),
565 "call_graph" => Some(CodeRetrievalLayer::CallGraph),
566 "import_graph" => Some(CodeRetrievalLayer::ImportGraph),
567 "sbom" => Some(CodeRetrievalLayer::Sbom),
568 "impact" => Some(CodeRetrievalLayer::Impact),
569 "text_fallback" => Some(CodeRetrievalLayer::TextFallback),
570 _ => None,
571 }
572}
573
574fn serialized_context_bytes<T: serde::Serialize>(value: &T) -> usize {
575 serde_json::to_vec(value)
576 .map(|bytes| bytes.len())
577 .unwrap_or(usize::MAX / 4)
578}
579
580fn context_paths(pack: &CodeGraphContextPack) -> Vec<String> {
581 pack.entry_points
582 .iter()
583 .chain(&pack.related_symbols)
584 .chain(&pack.graph_paths)
585 .map(|hit| hit.path.clone())
586 .collect::<BTreeSet<_>>()
587 .into_iter()
588 .collect()
589}
590
591#[cfg(test)]
592#[path = "mod_tests.rs"]
593mod tests;