Skip to main content

locus_gateway/
orchestration.rs

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}