Skip to main content

locus_mcp/
composition.rs

1use std::sync::Arc;
2
3use anyhow::Result;
4use locus_core_rs::domain::contracts::EmbeddingProvider;
5use locus_core_rs::{ParseProfile, SurrealDbSettings};
6#[cfg(feature = "local-embedding")]
7use locus_sdk::infrastructure::embeddings::LocalEmbeddingProvider;
8use locus_sdk::infrastructure::embeddings::OllamaEmbeddingProvider;
9use tracing::info;
10
11pub(crate) use locus_surreal_adapter::RuntimeSurrealDbClient;
12
13#[derive(Debug, Clone)]
14enum EmbeddingsProviderKind {
15    Ollama,
16    #[cfg(feature = "local-embedding")]
17    Local,
18}
19
20impl EmbeddingsProviderKind {
21    fn parse(value: &str) -> Option<Self> {
22        match value.trim().to_ascii_lowercase().as_str() {
23            "ollama" => Some(Self::Ollama),
24            #[cfg(feature = "local-embedding")]
25            "local" | "local-embedding" | "candle" => Some(Self::Local),
26            _ => None,
27        }
28    }
29}
30
31pub(crate) fn init_logging() {
32    let _ = tracing_subscriber::fmt()
33        .with_writer(std::io::stderr)
34        .with_env_filter(
35            tracing_subscriber::EnvFilter::try_from_default_env()
36                .unwrap_or_else(|_| tracing_subscriber::EnvFilter::new("info")),
37        )
38        .try_init();
39}
40
41pub(crate) fn load_surreal_settings(args: &[String]) -> Result<SurrealDbSettings> {
42    let mut settings = SurrealDbSettings::default();
43    settings.endpoints.embedded = Some("surrealkv://data/locus-mcp".to_string());
44    settings.database = "locus_mcp".to_string();
45
46    if let Some(value) = env_or_arg(
47        "LOCUS_MCP_SURREAL_REMOTE_ENDPOINT",
48        args,
49        "--remote-endpoint",
50    ) {
51        settings.endpoints.remote = Some(value);
52    }
53    if let Some(value) = env_or_arg(
54        "LOCUS_MCP_SURREAL_EMBEDDED_ENDPOINT",
55        args,
56        "--embedded-endpoint",
57    ) {
58        settings.endpoints.embedded = Some(value);
59    }
60    if let Some(value) = env_or_arg("LOCUS_MCP_SURREAL_ENDPOINT", args, "--endpoint") {
61        settings.endpoints.remote = Some(value.clone());
62        settings.endpoints.embedded = Some(value);
63    }
64    if let Some(value) = env_or_arg("LOCUS_MCP_SURREAL_NAMESPACE", args, "--namespace") {
65        settings.namespace = value;
66    }
67    if let Some(value) = env_or_arg("LOCUS_MCP_SURREAL_DATABASE", args, "--database") {
68        settings.database = value;
69    }
70    if let Some(value) = env_or_arg("LOCUS_MCP_SURREAL_USERNAME", args, "--username") {
71        settings.user = Some(value);
72    }
73    if let Some(value) = env_or_arg("LOCUS_MCP_SURREAL_PASSWORD", args, "--password") {
74        settings.password = Some(value);
75    }
76
77    Ok(settings)
78}
79
80pub(crate) fn runtime_args(args: &[String]) -> Vec<String> {
81    let mut runtime_args = args.to_vec();
82    if env_flag("LOCUS_MCP_REMOTE") && !runtime_args.iter().any(|value| value == "--remote") {
83        runtime_args.push("--remote".to_string());
84    }
85    runtime_args
86}
87
88pub(crate) fn build_embedding_provider(args: &[String]) -> Result<Option<Arc<dyn EmbeddingProvider>>> {
89    let embeddings_enabled = env_flag("LOCUS_MCP_EMBEDDINGS_ENABLED")
90        || args
91            .iter()
92            .any(|arg| arg.eq_ignore_ascii_case("--embeddings-enabled"));
93
94    if !embeddings_enabled {
95        return Ok(None);
96    }
97
98    let provider_kind_raw = env_or_arg(
99        "LOCUS_MCP_EMBEDDINGS_PROVIDER",
100        args,
101        "--embeddings-provider",
102    )
103    .unwrap_or_else(|| "ollama".to_string());
104    let provider_kind = EmbeddingsProviderKind::parse(&provider_kind_raw).ok_or_else(|| {
105        anyhow::anyhow!(
106            "unsupported embeddings provider '{}'; expected 'ollama'{}",
107            provider_kind_raw,
108            if cfg!(feature = "local-embedding") {
109                " or 'local'"
110            } else {
111                ""
112            }
113        )
114    })?;
115
116    let endpoint = env_or_arg(
117        "LOCUS_MCP_EMBEDDINGS_ENDPOINT",
118        args,
119        "--embeddings-endpoint",
120    )
121    .unwrap_or_else(|| "http://127.0.0.1:11434/api/embeddings".to_string());
122    let model = env_or_arg("LOCUS_MCP_EMBEDDINGS_MODEL", args, "--embeddings-model")
123        .unwrap_or_else(|| "sttp-encoder".to_string());
124    #[cfg(feature = "local-embedding")]
125    let repo = env_or_arg("LOCUS_MCP_EMBEDDINGS_REPO", args, "--embeddings-repo")
126        .unwrap_or_else(|| "sentence-transformers/all-MiniLM-L6-v2".to_string());
127
128    let provider: Arc<dyn EmbeddingProvider> = match provider_kind {
129        EmbeddingsProviderKind::Ollama => {
130            info!(
131                provider = "ollama",
132                endpoint = %endpoint,
133                model = %model,
134                "auto-embedding enabled for store_context"
135            );
136            Arc::new(OllamaEmbeddingProvider::new(endpoint, model))
137        }
138        #[cfg(feature = "local-embedding")]
139        EmbeddingsProviderKind::Local => {
140            info!(
141                provider = "local",
142                model = %model,
143                repo = %repo,
144                "auto-embedding enabled for store_context"
145            );
146            Arc::new(LocalEmbeddingProvider::new(model, repo)?)
147        }
148    };
149
150    Ok(Some(provider))
151}
152
153pub(crate) fn resolve_parser_profile(args: &[String]) -> Result<ParseProfile> {
154    let raw = env_or_arg("LOCUS_MCP_PARSE_PROFILE", args, "--parse-profile")
155        .unwrap_or_else(|| "strict_typed_ir".to_string());
156
157    parse_profile(raw.as_str()).ok_or_else(|| {
158        anyhow::anyhow!(
159            "unsupported parse profile '{}'; expected one of: strict_typed_ir, strict, tolerant",
160            raw
161        )
162    })
163}
164
165fn parse_profile(value: &str) -> Option<ParseProfile> {
166    match value.trim().to_ascii_lowercase().as_str() {
167        "strict_typed_ir" | "strict-typed-ir" | "stricttypedir" | "typed_ir" | "typed-ir" => {
168            Some(ParseProfile::StrictTypedIr)
169        }
170        "strict" => Some(ParseProfile::Strict),
171        "tolerant" | "default" => Some(ParseProfile::Tolerant),
172        _ => None,
173    }
174}
175
176fn env_or_arg(env_key: &str, args: &[String], arg_name: &str) -> Option<String> {
177    if let Ok(value) = std::env::var(env_key) {
178        let trimmed = value.trim();
179        if !trimmed.is_empty() {
180            return Some(trimmed.to_string());
181        }
182    }
183
184    arg_value(args, arg_name)
185}
186
187fn arg_value(args: &[String], key: &str) -> Option<String> {
188    args.windows(2)
189        .find(|window| window[0].eq_ignore_ascii_case(key))
190        .map(|window| window[1].clone())
191}
192
193fn env_flag(key: &str) -> bool {
194    std::env::var(key)
195        .map(|value| {
196            let normalized = value.trim().to_ascii_lowercase();
197            normalized == "1" || normalized == "true" || normalized == "yes"
198        })
199        .unwrap_or(false)
200}