relay_knowledge/application/runtime/
retrieval.rs1use std::{error::Error, fmt, time::Duration};
2
3use crate::{
4 domain::{RerankMode, RerankModeError},
5 env::{
6 RELAY_KNOWLEDGE_EMBEDDING_API_KEY, RELAY_KNOWLEDGE_EMBEDDING_BASE_URL,
7 RELAY_KNOWLEDGE_EMBEDDING_DIMENSION, RELAY_KNOWLEDGE_IMAGE_EMBEDDING_MODEL,
8 RELAY_KNOWLEDGE_RERANK_MODEL, RELAY_KNOWLEDGE_TEXT_EMBEDDING_MODEL, RetrievalEnvOverrides,
9 },
10 retrieval::{
11 DEFAULT_EMBEDDING_BATCH_SIZE, DEFAULT_EMBEDDING_MAX_CONCURRENCY, DEFAULT_EMBEDDING_TIMEOUT,
12 DEFAULT_RERANK_CANDIDATE_MULTIPLIER, DEFAULT_RERANK_MAX_CANDIDATES, DEFAULT_RERANK_TIMEOUT,
13 EmbeddingProviderKind, EmbeddingProviderKindError, LOCAL_RERANK_MODEL,
14 LOCAL_SEMANTIC_MODEL, LOCAL_VECTOR_DIMENSION, LOCAL_VECTOR_MODEL, ReadModelBackendConfig,
15 ReadModelBackendMode, ReadModelBackendModeError, ReadModelMetadata, RemoteEmbeddingConfig,
16 RerankConfig,
17 },
18};
19
20#[derive(Debug, Clone, PartialEq, Eq)]
22pub enum RetrievalRuntimeConfigError {
23 InvalidBackend(ReadModelBackendModeError),
24 InvalidRerankBackend(RerankModeError),
25 InvalidProvider(EmbeddingProviderKindError),
26 EmptyModelName(&'static str),
27 MissingRemoteValue(&'static str),
28 InvalidRemoteBaseUrl(String),
29 DimensionTooLarge(usize),
30}
31
32impl fmt::Display for RetrievalRuntimeConfigError {
33 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
34 match self {
35 Self::InvalidBackend(error) => write!(formatter, "{error}"),
36 Self::InvalidRerankBackend(error) => write!(formatter, "{error}"),
37 Self::InvalidProvider(error) => write!(formatter, "{error}"),
38 Self::EmptyModelName(variable) => {
39 write!(formatter, "{variable} must not be blank")
40 }
41 Self::MissingRemoteValue(variable) => {
42 write!(
43 formatter,
44 "{variable} is required when a read model backend is external"
45 )
46 }
47 Self::InvalidRemoteBaseUrl(value) => {
48 write!(
49 formatter,
50 "embedding base URL '{value}' must use http:// or https://"
51 )
52 }
53 Self::DimensionTooLarge(value) => {
54 write!(formatter, "embedding dimension {value} does not fit in u32")
55 }
56 }
57 }
58}
59
60impl Error for RetrievalRuntimeConfigError {}
61
62pub(super) fn retrieval_config_from_environment(
63 overrides: &RetrievalEnvOverrides,
64) -> Result<ReadModelBackendConfig, RetrievalRuntimeConfigError> {
65 let semantic_mode = parse_backend_mode(overrides.semantic_backend.as_deref())?;
66 let vector_mode = parse_backend_mode(overrides.vector_backend.as_deref())?;
67 let remote_required = semantic_mode == ReadModelBackendMode::External
68 || vector_mode == ReadModelBackendMode::External;
69 require_remote_model_metadata(overrides, remote_required)?;
70 let dimension = match overrides.embedding_dimension {
71 Some(value) => u32::try_from(value)
72 .map_err(|_| RetrievalRuntimeConfigError::DimensionTooLarge(value))?,
73 None => LOCAL_VECTOR_DIMENSION,
74 };
75 let text_model = model_name_override(
76 overrides.text_embedding_model.as_deref(),
77 RELAY_KNOWLEDGE_TEXT_EMBEDDING_MODEL,
78 LOCAL_VECTOR_MODEL,
79 )?;
80 let semantic_model = model_name_override(
81 overrides.text_embedding_model.as_deref(),
82 RELAY_KNOWLEDGE_TEXT_EMBEDDING_MODEL,
83 LOCAL_SEMANTIC_MODEL,
84 )?;
85 let image_model = model_name_override(
86 overrides.image_embedding_model.as_deref(),
87 RELAY_KNOWLEDGE_IMAGE_EMBEDDING_MODEL,
88 "relay-local-image-hash-v1",
89 )?;
90
91 let remote_embedding = remote_embedding_config_from_environment(overrides, remote_required)?;
92 let rerank = rerank_config_from_environment(overrides)?;
93
94 Ok(ReadModelBackendConfig {
95 semantic_mode,
96 vector_mode,
97 semantic_model: ReadModelMetadata {
98 name: semantic_model,
99 dimension,
100 },
101 vector_model: ReadModelMetadata {
102 name: text_model,
103 dimension,
104 },
105 image_model: ReadModelMetadata {
106 name: image_model,
107 dimension,
108 },
109 remote_embedding,
110 rerank,
111 })
112}
113
114fn rerank_config_from_environment(
115 overrides: &RetrievalEnvOverrides,
116) -> Result<RerankConfig, RetrievalRuntimeConfigError> {
117 let mode = overrides
118 .rerank_backend
119 .as_deref()
120 .map(RerankMode::parse)
121 .transpose()
122 .map_err(RetrievalRuntimeConfigError::InvalidRerankBackend)?
123 .unwrap_or(RerankMode::Local);
124 let model = match mode {
125 RerankMode::Disabled => None,
126 RerankMode::Local => Some(model_name_override(
127 overrides.rerank_model.as_deref(),
128 RELAY_KNOWLEDGE_RERANK_MODEL,
129 LOCAL_RERANK_MODEL,
130 )?),
131 RerankMode::External => overrides
132 .rerank_model
133 .as_deref()
134 .map(|model| model_name_override(Some(model), RELAY_KNOWLEDGE_RERANK_MODEL, ""))
135 .transpose()?,
136 };
137 let timeout = overrides
138 .rerank_timeout_ms
139 .map(Duration::from_millis)
140 .unwrap_or(DEFAULT_RERANK_TIMEOUT);
141
142 Ok(RerankConfig {
143 mode,
144 model,
145 timeout,
146 candidate_multiplier: overrides
147 .rerank_candidate_multiplier
148 .unwrap_or(DEFAULT_RERANK_CANDIDATE_MULTIPLIER),
149 max_candidates: overrides
150 .rerank_max_candidates
151 .unwrap_or(DEFAULT_RERANK_MAX_CANDIDATES),
152 })
153}
154
155fn require_remote_model_metadata(
156 overrides: &RetrievalEnvOverrides,
157 required: bool,
158) -> Result<(), RetrievalRuntimeConfigError> {
159 if !required {
160 return Ok(());
161 }
162 if overrides.text_embedding_model.is_none() {
163 return Err(RetrievalRuntimeConfigError::MissingRemoteValue(
164 RELAY_KNOWLEDGE_TEXT_EMBEDDING_MODEL,
165 ));
166 }
167 if overrides.embedding_dimension.is_none() {
168 return Err(RetrievalRuntimeConfigError::MissingRemoteValue(
169 RELAY_KNOWLEDGE_EMBEDDING_DIMENSION,
170 ));
171 }
172
173 Ok(())
174}
175
176fn remote_embedding_config_from_environment(
177 overrides: &RetrievalEnvOverrides,
178 required: bool,
179) -> Result<Option<RemoteEmbeddingConfig>, RetrievalRuntimeConfigError> {
180 if !required {
181 return Ok(None);
182 }
183 let provider = overrides
184 .llm_provider
185 .as_deref()
186 .map(EmbeddingProviderKind::parse)
187 .transpose()
188 .map_err(RetrievalRuntimeConfigError::InvalidProvider)?
189 .unwrap_or(EmbeddingProviderKind::OpenAiCompatible);
190 let base_url = required_remote_value(
191 overrides.embedding_base_url.as_deref(),
192 RELAY_KNOWLEDGE_EMBEDDING_BASE_URL,
193 )?;
194 if !base_url.starts_with("http://") && !base_url.starts_with("https://") {
195 return Err(RetrievalRuntimeConfigError::InvalidRemoteBaseUrl(base_url));
196 }
197 let api_key = required_remote_value(
198 overrides.embedding_api_key.as_deref(),
199 RELAY_KNOWLEDGE_EMBEDDING_API_KEY,
200 )?;
201 let batch_size = overrides
202 .embedding_batch_size
203 .unwrap_or(DEFAULT_EMBEDDING_BATCH_SIZE);
204 let timeout = overrides
205 .embedding_timeout_ms
206 .map(Duration::from_millis)
207 .unwrap_or(DEFAULT_EMBEDDING_TIMEOUT);
208 let max_concurrency = overrides
209 .embedding_max_concurrency
210 .unwrap_or(DEFAULT_EMBEDDING_MAX_CONCURRENCY);
211
212 Ok(Some(RemoteEmbeddingConfig {
213 provider,
214 base_url,
215 api_key,
216 batch_size,
217 timeout,
218 max_concurrency,
219 }))
220}
221
222fn required_remote_value(
223 value: Option<&str>,
224 variable: &'static str,
225) -> Result<String, RetrievalRuntimeConfigError> {
226 match value.map(str::trim) {
227 Some(trimmed) if !trimmed.is_empty() => Ok(trimmed.to_owned()),
228 _ => Err(RetrievalRuntimeConfigError::MissingRemoteValue(variable)),
229 }
230}
231
232fn model_name_override(
233 value: Option<&str>,
234 variable: &'static str,
235 default: &'static str,
236) -> Result<String, RetrievalRuntimeConfigError> {
237 match value {
238 Some(raw) => {
239 let trimmed = raw.trim();
240 if trimmed.is_empty() {
241 Err(RetrievalRuntimeConfigError::EmptyModelName(variable))
242 } else {
243 Ok(trimmed.to_owned())
244 }
245 }
246 None => Ok(default.to_owned()),
247 }
248}
249
250fn parse_backend_mode(
251 value: Option<&str>,
252) -> Result<ReadModelBackendMode, RetrievalRuntimeConfigError> {
253 value
254 .map(ReadModelBackendMode::parse)
255 .transpose()
256 .map_err(RetrievalRuntimeConfigError::InvalidBackend)
257 .map(|mode| mode.unwrap_or(ReadModelBackendMode::Local))
258}
259
260#[cfg(test)]
261#[path = "retrieval_tests.rs"]
262mod tests;