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