relay_knowledge/domain/code/repository/
retrieval_request.rs1use serde::{Deserialize, Serialize};
2
3use super::super::{DomainError, FreshnessPolicy, error::required_text};
4use super::CodeRepositorySelector;
5
6#[derive(Debug, Clone, PartialEq, Eq)]
7struct FieldQualifiers {
8 search_text: String,
9 kind_filters: Vec<String>,
10 language_filters: Vec<String>,
11 path_substrings: Vec<String>,
12 name_substrings: Vec<String>,
13}
14
15fn parse_field_qualifiers(query: &str) -> FieldQualifiers {
16 let mut plain_terms = Vec::new();
17 let mut qualifiers = FieldQualifiers {
18 search_text: String::new(),
19 kind_filters: Vec::new(),
20 language_filters: Vec::new(),
21 path_substrings: Vec::new(),
22 name_substrings: Vec::new(),
23 };
24
25 for token in query.split_whitespace() {
26 if !push_field_qualifier(token, &mut qualifiers) {
27 plain_terms.push(token);
28 }
29 }
30
31 qualifiers.search_text = plain_terms.join(" ");
32 if qualifiers.search_text.is_empty() && !query.trim().is_empty() {
33 qualifiers.search_text = query.trim().to_owned();
34 }
35
36 qualifiers
37}
38
39fn push_field_qualifier(token: &str, qualifiers: &mut FieldQualifiers) -> bool {
40 let Some((prefix, value)) = token.split_once(':') else {
41 return false;
42 };
43 if value.trim().is_empty() || value.starts_with(':') {
44 return false;
45 }
46
47 match prefix.to_ascii_lowercase().as_str() {
48 "kind" => {
49 extend_qualifier_values(&mut qualifiers.kind_filters, value, true);
50 true
51 }
52 "lang" | "language" => {
53 extend_qualifier_values(&mut qualifiers.language_filters, value, true);
54 true
55 }
56 "path" => {
57 extend_qualifier_values(&mut qualifiers.path_substrings, value, false);
58 true
59 }
60 "name" => {
61 extend_qualifier_values(&mut qualifiers.name_substrings, value, false);
62 true
63 }
64 _ => false,
65 }
66}
67
68fn extend_qualifier_values(values: &mut Vec<String>, raw_value: &str, ascii_lowercase: bool) {
69 for value in raw_value
70 .split(',')
71 .map(str::trim)
72 .filter(|value| !value.is_empty())
73 {
74 let value = if ascii_lowercase {
75 value.to_ascii_lowercase()
76 } else {
77 value.to_owned()
78 };
79 if !values.contains(&value) {
80 values.push(value);
81 }
82 }
83}
84
85#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
87#[serde(rename_all = "snake_case")]
88pub enum CodeQueryKind {
89 Hybrid,
90 Symbol,
91 Definition,
92 References,
93 Callers,
94 Callees,
95 Imports,
96 Sbom,
97 Impact,
98}
99
100#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
102pub struct CodeRetrievalRequest {
103 pub query: String,
104 pub repository: CodeRepositorySelector,
105 pub code_query_kind: CodeQueryKind,
106 pub limit: usize,
107 pub freshness_policy: FreshnessPolicy,
108 #[serde(default)]
109 pub exclude_generated: bool,
110 #[serde(default, skip_serializing_if = "Vec::is_empty")]
111 pub query_kind_filters: Vec<String>,
112 #[serde(default, skip_serializing_if = "Vec::is_empty")]
113 pub query_language_filters: Vec<String>,
114 #[serde(default, skip_serializing_if = "Vec::is_empty")]
115 pub query_path_substrings: Vec<String>,
116 #[serde(default, skip_serializing_if = "Vec::is_empty")]
117 pub query_name_substrings: Vec<String>,
118}
119
120impl CodeRetrievalRequest {
121 pub fn new(
123 query: impl Into<String>,
124 repository: CodeRepositorySelector,
125 code_query_kind: CodeQueryKind,
126 limit: usize,
127 freshness_policy: FreshnessPolicy,
128 ) -> Result<Self, DomainError> {
129 let limit = match limit {
130 1..=50 => limit,
131 0 => return Err(DomainError::invalid("limit", "must be greater than zero")),
132 _ => return Err(DomainError::invalid("limit", "must be 50 or less")),
133 };
134
135 let qualifiers = parse_field_qualifiers(&required_text("query", query)?);
136
137 Ok(Self {
138 query: qualifiers.search_text,
139 repository,
140 code_query_kind,
141 limit,
142 freshness_policy,
143 exclude_generated: false,
144 query_kind_filters: qualifiers.kind_filters,
145 query_language_filters: qualifiers.language_filters,
146 query_path_substrings: qualifiers.path_substrings,
147 query_name_substrings: qualifiers.name_substrings,
148 })
149 }
150}
151
152#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
154pub struct CodeFeatureFlagRequest {
155 #[serde(skip_serializing_if = "Option::is_none")]
156 pub query: Option<String>,
157 pub repository: CodeRepositorySelector,
158 pub limit: usize,
159 pub freshness_policy: FreshnessPolicy,
160}
161
162impl CodeFeatureFlagRequest {
163 pub fn new(
165 query: Option<String>,
166 repository: CodeRepositorySelector,
167 limit: usize,
168 freshness_policy: FreshnessPolicy,
169 ) -> Result<Self, DomainError> {
170 let limit = match limit {
171 1..=100 => limit,
172 0 => return Err(DomainError::invalid("limit", "must be greater than zero")),
173 _ => return Err(DomainError::invalid("limit", "must be 100 or less")),
174 };
175 let query = query
176 .map(|value| required_text("query", value))
177 .transpose()?;
178
179 Ok(Self {
180 query,
181 repository,
182 limit,
183 freshness_policy,
184 })
185 }
186}
187
188#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
190pub struct CodeImpactRequest {
191 pub repository: CodeRepositorySelector,
192 pub base_ref: String,
193 pub head_ref: String,
194 pub limit: usize,
195}
196
197impl CodeImpactRequest {
198 pub fn new(
200 repository: CodeRepositorySelector,
201 base_ref: impl Into<String>,
202 head_ref: impl Into<String>,
203 limit: usize,
204 ) -> Result<Self, DomainError> {
205 let limit = match limit {
206 1..=100 => limit,
207 0 => return Err(DomainError::invalid("limit", "must be greater than zero")),
208 _ => return Err(DomainError::invalid("limit", "must be 100 or less")),
209 };
210
211 Ok(Self {
212 repository,
213 base_ref: required_text("base_ref", base_ref)?,
214 head_ref: required_text("head_ref", head_ref)?,
215 limit,
216 })
217 }
218}
219
220#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
222#[serde(rename_all = "snake_case")]
223pub enum CodeRetrievalLayer {
224 Lexical,
225 Symbol,
226 Definition,
227 Reference,
228 CallGraph,
229 ImportGraph,
230 Sbom,
231 Impact,
232 TextFallback,
233}
234
235impl CodeRetrievalLayer {
236 pub const fn as_str(self) -> &'static str {
238 match self {
239 Self::Lexical => "lexical",
240 Self::Symbol => "symbol",
241 Self::Definition => "definition",
242 Self::Reference => "reference",
243 Self::CallGraph => "call_graph",
244 Self::ImportGraph => "import_graph",
245 Self::Sbom => "sbom",
246 Self::Impact => "impact",
247 Self::TextFallback => "text_fallback",
248 }
249 }
250}
251
252#[cfg(test)]
253#[path = "retrieval_request_tests.rs"]
254mod tests;