1use std::sync::Arc;
2
3use anyhow::{Result, anyhow};
4use axum::http::HeaderValue;
5use locus_core_rs::application::services::{
6 CalibrationService, ContextQueryService, MonthlyRollupService,
7 MoodCatalogService, RekeyScopeService, StoreContextService,
8};
9use locus_core_rs::application::validation::TreeSitterValidator;
10use locus_core_rs::domain::contracts::{
11 EmbeddingProvider, NodeStore, NodeStoreInitializer, NodeValidator, SemanticIndexStore,
12 SemanticIndexStoreInitializer,
13};
14use locus_core_rs::parsing::SttpNodeParser;
15use locus_core_rs::storage::{
16 InMemoryNodeStore, InMemorySemanticIndexStore, SurrealDbEndpointsSettings, SurrealDbNodeStore,
17 SurrealDbRuntimeOptions, SurrealDbSemanticIndexStore, SurrealDbSettings,
18};
19#[cfg(feature = "local-embedding")]
20use locus_sdk::infrastructure::embeddings::LocalEmbeddingProvider;
21use locus_sdk::infrastructure::embeddings::OllamaEmbeddingProvider;
22use tracing::{error, info};
23
24use crate::app_state::AppState;
25use crate::gateway_args::{EmbeddingsProviderKind, GatewayArgs, GatewayBackend};
26use crate::http_models::CorsAllowedOrigins;
27use crate::providers::{AvecScorer, OllamaAvecScorer};
28use crate::surreal_client::RuntimeSurrealDbClient;
29
30pub(crate) async fn build_state(args: &GatewayArgs) -> Result<AppState> {
31 build_state_with_backend(&args.backend, Some(args)).await
32}
33
34#[cfg(test)]
35pub(crate) async fn build_in_memory_state() -> Result<AppState> {
36 build_in_memory_state_with_args(None).await
37}
38
39pub(crate) fn parse_cors_allowed_origins(value: &str) -> Result<CorsAllowedOrigins> {
40 let trimmed = value.trim();
41 if trimmed.is_empty() {
42 return Err(anyhow!(
43 "CORS allowed origins cannot be empty when CORS is enabled"
44 ));
45 }
46
47 if trimmed == "*" {
48 return Ok(CorsAllowedOrigins::Any);
49 }
50
51 let mut origins = Vec::new();
52 for origin in trimmed
53 .split(',')
54 .map(str::trim)
55 .filter(|part| !part.is_empty())
56 {
57 let header = HeaderValue::from_str(origin)
58 .map_err(|_| anyhow!("Invalid CORS origin value: {origin}"))?;
59 origins.push(header);
60 }
61
62 if origins.is_empty() {
63 return Err(anyhow!(
64 "CORS allowed origins must include at least one origin or '*'"
65 ));
66 }
67
68 Ok(CorsAllowedOrigins::Explicit(origins))
69}
70
71pub(crate) async fn shutdown_signal() {
72 if let Err(err) = tokio::signal::ctrl_c().await {
73 error!(error = %err, "Failed waiting for ctrl_c signal");
74 }
75}
76
77async fn build_state_with_backend(
78 backend: &GatewayBackend,
79 options: Option<&GatewayArgs>,
80) -> Result<AppState> {
81 match backend {
82 GatewayBackend::InMemory => build_in_memory_state_with_args(options).await,
83 GatewayBackend::Surreal => {
84 let options = options.ok_or_else(|| {
85 anyhow!("Surreal backend selected, but no gateway runtime options were provided.")
86 })?;
87 build_surreal_state(options).await
88 }
89 }
90}
91
92async fn build_in_memory_state_with_args(args: Option<&GatewayArgs>) -> Result<AppState> {
93 let store = Arc::new(InMemoryNodeStore::new());
94 let semantic_index = Arc::new(InMemorySemanticIndexStore::new());
95
96 let initializer: Arc<dyn NodeStoreInitializer> = store.clone();
97 initializer.initialize_async().await?;
98
99 let semantic_initializer: Arc<dyn SemanticIndexStoreInitializer> = semantic_index.clone();
100 semantic_initializer.initialize_async().await?;
101
102 let store_trait: Arc<dyn NodeStore> = store;
103 let semantic_trait: Arc<dyn SemanticIndexStore> = semantic_index;
104 let validator: Arc<dyn NodeValidator> = Arc::new(TreeSitterValidator);
105 let embedding_provider = build_embedding_provider(args)?;
106 let avec_scorer = build_avec_scorer(args);
107
108 Ok(build_services(
109 store_trait,
110 semantic_trait,
111 validator,
112 embedding_provider,
113 avec_scorer,
114 ))
115}
116
117fn build_services(
118 store_trait: Arc<dyn NodeStore>,
119 semantic_index: Arc<dyn SemanticIndexStore>,
120 validator: Arc<dyn NodeValidator>,
121 embedding_provider: Option<Arc<dyn EmbeddingProvider>>,
122 avec_scorer: Option<Arc<dyn AvecScorer>>,
123) -> AppState {
124 let parser = SttpNodeParser::new();
125 let store_context = match embedding_provider.as_ref() {
126 Some(provider) => Arc::new(
127 StoreContextService::with_embedding_provider(
128 store_trait.clone(),
129 validator.clone(),
130 provider.clone(),
131 parser,
132 )
133 .with_semantic_index(semantic_index.clone()),
134 ),
135 None => Arc::new(
136 StoreContextService::new(store_trait.clone(), validator.clone(), SttpNodeParser::new())
137 .with_semantic_index(semantic_index.clone()),
138 ),
139 };
140
141 let mut monthly_rollup =
142 MonthlyRollupService::new(store_trait.clone(), validator.clone())
143 .with_semantic_index(semantic_index.clone());
144 if let Some(provider) = embedding_provider.as_ref() {
145 monthly_rollup = monthly_rollup.with_embedding_provider(provider.clone());
146 }
147
148 AppState {
149 node_store: store_trait.clone(),
150 semantic_index,
151 embedding_provider: embedding_provider.clone(),
152 avec_scorer,
153 calibration: Arc::new(CalibrationService::new(store_trait.clone())),
154 context_query: Arc::new(ContextQueryService::new(store_trait.clone())),
155 mood_catalog: Arc::new(MoodCatalogService::new()),
156 store_context,
157 monthly_rollup: Arc::new(monthly_rollup),
158 rekey_scope: Arc::new(RekeyScopeService::new(store_trait)),
159 }
160}
161
162async fn build_surreal_state(args: &GatewayArgs) -> Result<AppState> {
163 let mut settings = SurrealDbSettings::default();
164 settings.endpoints = SurrealDbEndpointsSettings {
165 embedded: args
166 .surreal_embedded_endpoint
167 .clone()
168 .or(settings.endpoints.embedded),
169 remote: args
170 .surreal_remote_endpoint
171 .clone()
172 .or(settings.endpoints.remote),
173 };
174 settings.namespace = args.surreal_namespace.clone();
175 settings.database = args.surreal_database.clone();
176 settings.user = Some(args.surreal_user.clone());
177 settings.password = Some(args.surreal_password.clone());
178
179 let mut runtime_args = Vec::new();
180 if args.remote {
181 runtime_args.push("--remote".to_string());
182 }
183
184 let runtime = SurrealDbRuntimeOptions::from_args(
185 &runtime_args,
186 &settings,
187 Some(args.root_dir_name.as_str()),
188 )?;
189
190 info!(
191 backend = "surreal",
192 root_dir = runtime.root_dir,
193 mode = if runtime.use_remote {
194 "remote"
195 } else {
196 "embedded"
197 },
198 endpoint = runtime.endpoint,
199 namespace = runtime.namespace,
200 database = runtime.database,
201 "Surreal backend requested"
202 );
203
204 let client = Arc::new(
205 RuntimeSurrealDbClient::connect(
206 &runtime,
207 settings.user.as_deref(),
208 settings.password.as_deref(),
209 )
210 .await?,
211 );
212
213 let semantic_index = Arc::new(SurrealDbSemanticIndexStore::new(client.clone()));
214 let store = Arc::new(SurrealDbNodeStore::new(client));
215
216 let initializer: Arc<dyn NodeStoreInitializer> = store.clone();
217 initializer.initialize_async().await?;
218
219 let semantic_initializer: Arc<dyn SemanticIndexStoreInitializer> = semantic_index.clone();
220 semantic_initializer.initialize_async().await?;
221
222 let store_trait: Arc<dyn NodeStore> = store;
223 let semantic_trait: Arc<dyn SemanticIndexStore> = semantic_index;
224 let validator: Arc<dyn NodeValidator> = Arc::new(TreeSitterValidator);
225 let embedding_provider = build_embedding_provider(Some(args))?;
226 let avec_scorer = build_avec_scorer(Some(args));
227
228 Ok(build_services(
229 store_trait,
230 semantic_trait,
231 validator,
232 embedding_provider,
233 avec_scorer,
234 ))
235}
236
237fn build_avec_scorer(args: Option<&GatewayArgs>) -> Option<Arc<dyn AvecScorer>> {
238 let args = args?;
239 if !args.avec_scoring_enabled {
240 return None;
241 }
242
243 Some(Arc::new(OllamaAvecScorer::new(
244 args.avec_scoring_endpoint.clone(),
245 args.avec_scoring_model.clone(),
246 )))
247}
248
249fn build_embedding_provider(
250 args: Option<&GatewayArgs>,
251) -> Result<Option<Arc<dyn EmbeddingProvider>>> {
252 let Some(args) = args else {
253 return Ok(None);
254 };
255
256 if !args.embeddings_enabled {
257 return Ok(None);
258 }
259
260 let provider: Arc<dyn EmbeddingProvider> = match args.embeddings_provider {
261 EmbeddingsProviderKind::Ollama => Arc::new(OllamaEmbeddingProvider::new(
262 args.embeddings_endpoint.clone(),
263 args.embeddings_model.clone(),
264 )),
265 #[cfg(feature = "local-embedding")]
266 EmbeddingsProviderKind::Local => Arc::new(LocalEmbeddingProvider::new(
267 args.embeddings_model.clone(),
268 args.embeddings_repo.clone(),
269 )?),
270 };
271
272 Ok(Some(provider))
273}