Skip to main content

cognee_http_server/
wiring.rs

1//! Construct default standalone backend handles for the HTTP server binary.
2
3use std::path::{Path, PathBuf};
4use std::sync::Arc;
5
6use anyhow::anyhow;
7use cognee_core::{CpuPool, RayonThreadPool};
8use cognee_database::{
9    CheckpointStore, DatabaseConnection, DeleteDb, IngestDb, SeaOrmCheckpointStore,
10    SearchHistoryDb, connect, initialize,
11};
12use cognee_delete::DeleteService;
13use cognee_embedding::{EmbeddingConfig, EmbeddingEngine, EmbeddingProvider};
14use cognee_graph::{GraphDBTrait, LadybugAdapter};
15use cognee_llm::{Llm, OpenAIAdapter, OpenAIResponsesClient, ResponsesClient, Transcriber};
16use cognee_ontology::{OntologyManager, OntologyResolver};
17use cognee_search::{
18    SeaOrmSessionStore, SearchBuilder, SearchOrchestrator, SessionManager, SessionStore,
19};
20use cognee_storage::{LocalStorage, StorageTrait};
21use cognee_vector::{PgVectorAdapter, VectorDB};
22use secrecy::ExposeSecret;
23
24use crate::components::ComponentHandles;
25use crate::config::HttpServerConfig;
26use crate::error::ServerError;
27use crate::notebook_runner::SubprocessRunner;
28
29fn ensure_dir(path: &Path) -> Result<(), ServerError> {
30    std::fs::create_dir_all(path)
31        .map_err(|e| ServerError::Other(anyhow!("create_dir_all({}): {e}", path.display())))
32}
33
34pub async fn wire_default_backends(
35    cfg: &HttpServerConfig,
36) -> Result<ComponentHandles, ServerError> {
37    ensure_dir(&cfg.data_root_directory)?;
38    ensure_dir(&cfg.system_root_directory)?;
39
40    let storage = wire_storage(cfg).await?;
41    let database = wire_database(cfg).await?;
42    let graph_db = wire_graph_db(cfg).await?;
43    let vector_db = wire_vector_db(cfg).await?;
44
45    let embedding_engine = wire_embedding_engine(cfg).await;
46    let llm = wire_llm(cfg);
47    let transcriber = wire_transcriber(cfg);
48
49    let thread_pool: Option<Arc<dyn CpuPool>> = Some(Arc::new(
50        RayonThreadPool::with_default_threads()
51            .map_err(|e| ServerError::Other(anyhow!("rayon thread pool init failed: {e}")))?,
52    ));
53
54    let ontology_manager = Arc::new(OntologyManager::new(
55        cfg.data_root_directory.join("ontology"),
56    ));
57    let ontology_resolver: Option<Arc<dyn OntologyResolver>> = None;
58
59    let delete_service = Arc::new(DeleteService::new(
60        Arc::clone(&storage),
61        Arc::clone(&database) as Arc<dyn DeleteDb>,
62    ));
63
64    let checkpoint_store = Some(
65        Arc::new(SeaOrmCheckpointStore::new(Arc::clone(&database))) as Arc<dyn CheckpointStore>
66    );
67
68    let (session_store, session_manager) = wire_session(cfg, Arc::clone(&database)).await;
69
70    let search_orchestrator = wire_search_orchestrator(
71        Arc::clone(&database),
72        llm.clone(),
73        Arc::clone(&graph_db),
74        Arc::clone(&vector_db),
75        embedding_engine.clone(),
76        session_manager.clone(),
77    );
78
79    let responses_client = wire_responses_client(cfg);
80
81    let notebook_runner = if cfg.notebook_runner_enabled {
82        Some(SubprocessRunner::new().into_dyn())
83    } else {
84        None
85    };
86
87    Ok(ComponentHandles {
88        database,
89        acl_db: None,
90        storage,
91        delete_service,
92        cloud_client: None,
93        ontology_manager,
94        search_orchestrator,
95        llm,
96        transcriber,
97        graph_db: Some(graph_db),
98        vector_db: Some(vector_db),
99        thread_pool,
100        embedding_engine,
101        ontology_resolver,
102        session_store,
103        session_manager,
104        checkpoint_store,
105        responses_client,
106        notebook_runner,
107    })
108}
109
110async fn wire_storage(cfg: &HttpServerConfig) -> Result<Arc<dyn StorageTrait>, ServerError> {
111    let storage =
112        Arc::new(LocalStorage::new(cfg.data_root_directory.clone())) as Arc<dyn StorageTrait>;
113    storage
114        .initialize()
115        .await
116        .map_err(|e| ServerError::Other(anyhow!("storage init failed: {e}")))?;
117    Ok(storage)
118}
119
120async fn wire_database(cfg: &HttpServerConfig) -> Result<Arc<DatabaseConnection>, ServerError> {
121    let url = cfg.relational_db_url.clone();
122
123    if let Some(path) = url.strip_prefix("sqlite://")
124        && !path.starts_with(':')
125    {
126        let db_path = PathBuf::from(path);
127        if let Some(parent) = db_path.parent() {
128            ensure_dir(parent)?;
129        }
130        if !db_path.exists() {
131            std::fs::File::create(&db_path).map_err(|e| {
132                ServerError::Other(anyhow!("create sqlite file {}: {e}", db_path.display()))
133            })?;
134        }
135    }
136
137    let db = connect(&url)
138        .await
139        .map_err(|e| ServerError::Other(anyhow!("database connect failed: {e}")))?;
140    initialize(&db)
141        .await
142        .map_err(|e| ServerError::Other(anyhow!("database migrate failed: {e}")))?;
143
144    Ok(Arc::new(db))
145}
146
147async fn wire_graph_db(cfg: &HttpServerConfig) -> Result<Arc<dyn GraphDBTrait>, ServerError> {
148    if !cfg.graph_provider.eq_ignore_ascii_case("ladybug") {
149        return Err(ServerError::Other(anyhow!(
150            "unsupported graph provider '{}'; only 'ladybug' is supported",
151            cfg.graph_provider
152        )));
153    }
154
155    if let Some(parent) = cfg.graph_file_path.parent() {
156        ensure_dir(parent)?;
157    }
158
159    let path = cfg.graph_file_path.to_string_lossy().to_string();
160    let graph = LadybugAdapter::new(&path)
161        .await
162        .map_err(|e| ServerError::Other(anyhow!("graph init failed: {e}")))?;
163    graph
164        .initialize()
165        .await
166        .map_err(|e| ServerError::Other(anyhow!("graph schema init failed: {e}")))?;
167    Ok(Arc::new(graph) as Arc<dyn GraphDBTrait>)
168}
169
170async fn wire_vector_db(cfg: &HttpServerConfig) -> Result<Arc<dyn VectorDB>, ServerError> {
171    // The qdrant adapter has been extracted to the closed `cognee-vector-qdrant`
172    // crate; the OSS http-server wires pgvector as the production default and
173    // exposes an opt-in in-memory mock behind the `dev-mock` feature so local
174    // dev / `cargo test` work without a Postgres instance.
175    let provider = cfg.vector_provider.to_ascii_lowercase();
176    match provider.as_str() {
177        "pgvector" => {
178            let url = cfg.vector_db_url.trim();
179            if url.is_empty() {
180                return Err(ServerError::Other(anyhow!(
181                    "VECTOR_DB_URL (postgres connection string) is required when \
182                     VECTOR_DB_PROVIDER=pgvector"
183                )));
184            }
185            let adapter = PgVectorAdapter::new(url, cfg.embedding_dimensions as usize)
186                .await
187                .map_err(|e| ServerError::Other(anyhow!("pgvector adapter init: {e}")))?;
188            Ok(Arc::new(adapter) as Arc<dyn VectorDB>)
189        }
190        #[cfg(feature = "dev-mock")]
191        "mock" => {
192            // OSS single-user dev path — keeps `cargo test` + local dev working
193            // without a Postgres instance. Off in production builds.
194            // The `dev-mock` feature enables `cognee-vector/testing`, which
195            // is where `MockVectorDB` actually lives.
196            Ok(Arc::new(cognee_vector::MockVectorDB::new()) as Arc<dyn VectorDB>)
197        }
198        other => Err(ServerError::Other(anyhow!(
199            "vector_db_provider='{other}' not supported in the OSS http-server. \
200             Supported: 'pgvector' (and 'mock' when built with the `dev-mock` \
201             feature). The Qdrant adapter has been extracted to the closed \
202             cognee-vector-qdrant crate."
203        ))),
204    }
205}
206
207fn build_embedding_config(cfg: &HttpServerConfig) -> Option<EmbeddingConfig> {
208    let provider = match cfg.embedding_provider.trim().to_ascii_lowercase().as_str() {
209        "onnx" => EmbeddingProvider::Onnx,
210        "fastembed" => EmbeddingProvider::Fastembed,
211        "openai" => EmbeddingProvider::OpenAi,
212        "openai_compatible" => EmbeddingProvider::OpenAiCompatible,
213        "ollama" => EmbeddingProvider::Ollama,
214        "mock" => EmbeddingProvider::Mock,
215        other => {
216            tracing::warn!("unknown embedding provider '{other}', embedding engine not wired");
217            return None;
218        }
219    };
220    let mut embedding_cfg = EmbeddingConfig {
221        provider,
222        model: cfg.embedding_model_name.clone(),
223        dimensions: cfg.embedding_dimensions as usize,
224        ..Default::default()
225    };
226    if !cfg.embedding_endpoint.trim().is_empty() {
227        embedding_cfg.endpoint = Some(cfg.embedding_endpoint.clone());
228    }
229    if !cfg.embedding_api_key.expose_secret().is_empty() {
230        embedding_cfg.api_key = Some(cfg.embedding_api_key.expose_secret().to_string());
231    }
232    embedding_cfg.onnx.model_name = cfg.embedding_model_name.clone();
233    embedding_cfg.onnx.dimensions = cfg.embedding_dimensions as usize;
234    if let Some(model_path) = &cfg.embedding_model_path {
235        embedding_cfg.onnx.model_path = model_path.clone();
236    }
237    if let Some(tokenizer_path) = &cfg.embedding_tokenizer_path {
238        embedding_cfg.onnx.tokenizer_path = tokenizer_path.clone();
239    }
240
241    Some(embedding_cfg)
242}
243
244async fn wire_embedding_engine(cfg: &HttpServerConfig) -> Option<Arc<dyn EmbeddingEngine>> {
245    let embedding_cfg = build_embedding_config(cfg)?;
246
247    match embedding_cfg.create_engine().await {
248        Ok(engine) => Some(engine),
249        Err(err) => {
250            tracing::warn!("embedding engine unavailable, wiring as None: {err}");
251            None
252        }
253    }
254}
255
256fn wire_llm(cfg: &HttpServerConfig) -> Option<Arc<dyn Llm>> {
257    if !cfg.llm_provider.eq_ignore_ascii_case("openai") {
258        tracing::warn!(
259            "LLM provider '{}' is not supported by standalone wiring yet; llm not wired",
260            cfg.llm_provider
261        );
262        return None;
263    }
264
265    let api_key = cfg.llm_api_key.expose_secret().to_string();
266    if api_key.is_empty() {
267        tracing::warn!("LLM API key missing; llm not wired");
268        return None;
269    }
270
271    let endpoint = if cfg.llm_endpoint.trim().is_empty() {
272        None
273    } else {
274        Some(cfg.llm_endpoint.clone())
275    };
276
277    match OpenAIAdapter::new(cfg.llm_model.clone(), api_key, endpoint).map(|adapter| {
278        adapter
279            .with_structured_output_retries(cfg.llm_max_retries.max(1))
280            .with_network_retries(cfg.llm_max_retries.max(1))
281    }) {
282        Ok(adapter) => Some(Arc::new(adapter) as Arc<dyn Llm>),
283        Err(err) => {
284            tracing::warn!("llm wiring failed, wiring as None: {err}");
285            None
286        }
287    }
288}
289
290fn wire_transcriber(cfg: &HttpServerConfig) -> Option<Arc<dyn Transcriber>> {
291    if !cfg.llm_provider.eq_ignore_ascii_case("openai") {
292        return None;
293    }
294
295    let api_key = cfg.llm_api_key.expose_secret().to_string();
296    if api_key.is_empty() {
297        return None;
298    }
299
300    let endpoint = if cfg.llm_endpoint.trim().is_empty() {
301        None
302    } else {
303        Some(cfg.llm_endpoint.clone())
304    };
305
306    match OpenAIAdapter::new(cfg.llm_model.clone(), api_key, endpoint).map(|adapter| {
307        adapter
308            .with_structured_output_retries(cfg.llm_max_retries.max(1))
309            .with_network_retries(cfg.llm_max_retries.max(1))
310    }) {
311        Ok(adapter) => Some(Arc::new(adapter) as Arc<dyn Transcriber>),
312        Err(err) => {
313            tracing::warn!("transcriber wiring failed, wiring as None: {err}");
314            None
315        }
316    }
317}
318
319async fn wire_session(
320    cfg: &HttpServerConfig,
321    database: Arc<DatabaseConnection>,
322) -> (Option<Arc<dyn SessionStore>>, Option<Arc<SessionManager>>) {
323    if !cfg.session_store_backend.eq_ignore_ascii_case("seaorm") {
324        tracing::warn!(
325            "session store backend '{}' unsupported in standalone wiring; session disabled",
326            cfg.session_store_backend
327        );
328        return (None, None);
329    }
330
331    match SeaOrmSessionStore::new(database).await {
332        Ok(store_impl) => {
333            let store: Arc<dyn SessionStore> = Arc::new(store_impl);
334            let manager = Arc::new(SessionManager::new(Arc::clone(&store)));
335            (Some(store), Some(manager))
336        }
337        Err(err) => {
338            tracing::warn!("session store wiring failed, wiring as None: {err}");
339            (None, None)
340        }
341    }
342}
343
344fn wire_search_orchestrator(
345    database: Arc<DatabaseConnection>,
346    llm: Option<Arc<dyn Llm>>,
347    graph_db: Arc<dyn GraphDBTrait>,
348    vector_db: Arc<dyn VectorDB>,
349    embedding_engine: Option<Arc<dyn EmbeddingEngine>>,
350    session_manager: Option<Arc<SessionManager>>,
351) -> Option<Arc<SearchOrchestrator>> {
352    let (Some(llm), Some(embedding_engine)) = (llm, embedding_engine) else {
353        tracing::warn!(
354            "search orchestrator not wired: requires llm + embedding engine, one or more missing"
355        );
356        return None;
357    };
358
359    let mut builder = SearchBuilder::new(
360        vector_db,
361        embedding_engine,
362        graph_db,
363        llm,
364        Arc::clone(&database) as Arc<dyn SearchHistoryDb>,
365    )
366    .with_dataset_resolver(Arc::clone(&database) as Arc<dyn IngestDb>);
367
368    if let Some(sm) = session_manager {
369        builder = builder.with_session_manager(sm);
370    }
371
372    Some(Arc::new(builder.build()))
373}
374
375fn wire_responses_client(cfg: &HttpServerConfig) -> Option<Arc<dyn ResponsesClient>> {
376    if !cfg.responses_client_enabled {
377        return None;
378    }
379
380    let api_key = cfg.llm_api_key.expose_secret().to_string();
381    if api_key.is_empty() {
382        tracing::warn!("responses client enabled but llm api key is missing; wiring as None");
383        return None;
384    }
385
386    let endpoint = if cfg.llm_endpoint.trim().is_empty() {
387        None
388    } else {
389        Some(cfg.llm_endpoint.clone())
390    };
391
392    match OpenAIResponsesClient::new(api_key, endpoint) {
393        Ok(client) => Some(Arc::new(client) as Arc<dyn ResponsesClient>),
394        Err(err) => {
395            tracing::warn!("responses client wiring failed, wiring as None: {err}");
396            None
397        }
398    }
399}
400
401#[cfg(test)]
402#[allow(
403    clippy::unwrap_used,
404    clippy::expect_used,
405    reason = "test code — panics are acceptable failures"
406)]
407mod tests {
408    use super::*;
409
410    #[test]
411    fn build_embedding_config_applies_explicit_onnx_asset_paths() {
412        let cfg = HttpServerConfig {
413            embedding_provider: "onnx".to_string(),
414            embedding_model_name: "custom-bge".to_string(),
415            embedding_dimensions: 768,
416            embedding_model_path: Some(PathBuf::from("/tmp/model.onnx")),
417            embedding_tokenizer_path: Some(PathBuf::from("/tmp/tokenizer.json")),
418            ..Default::default()
419        };
420
421        let embedding_cfg = build_embedding_config(&cfg).expect("embedding config");
422
423        assert_eq!(embedding_cfg.provider, EmbeddingProvider::Onnx);
424        assert_eq!(embedding_cfg.model, "custom-bge");
425        assert_eq!(embedding_cfg.dimensions, 768);
426        assert_eq!(embedding_cfg.onnx.model_name, "custom-bge");
427        assert_eq!(embedding_cfg.onnx.dimensions, 768);
428        assert_eq!(
429            embedding_cfg.onnx.model_path,
430            PathBuf::from("/tmp/model.onnx")
431        );
432        assert_eq!(
433            embedding_cfg.onnx.tokenizer_path,
434            PathBuf::from("/tmp/tokenizer.json")
435        );
436    }
437
438    #[tokio::test]
439    async fn wire_default_backends_fails_on_invalid_database_url() {
440        let mut cfg = HttpServerConfig::default();
441        let temp = tempfile::tempdir().expect("tempdir");
442        cfg.data_root_directory = temp.path().join("data");
443        cfg.system_root_directory = temp.path().join("system");
444        cfg.graph_file_path = cfg.system_root_directory.join("graph");
445        cfg.vector_db_url = cfg
446            .system_root_directory
447            .join("vectors")
448            .display()
449            .to_string();
450        cfg.relational_db_url = "not-a-valid-db-url".to_string();
451
452        let result = wire_default_backends(&cfg).await;
453        assert!(result.is_err());
454
455        let msg = match result {
456            Ok(_) => String::new(),
457            Err(err) => err.to_string(),
458        };
459        assert!(msg.contains("database connect failed"));
460    }
461}