Skip to main content

relay_knowledge/application/runtime/
retrieval.rs

1use 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/// Retrieval runtime configuration validation error.
21#[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;