Skip to main content

relay_knowledge/application/code_repository/context/
mod.rs

1//! Builds bounded graph-aware code context packs and provenance.
2
3use 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    /// Builds an agent-oriented codegraph context pack with bounded graph expansion.
37    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;