Skip to main content

soma_codemode/
local_provider.rs

1use serde_json::Value;
2
3use crate::git::provider::GitProvider;
4use crate::state::provider::StateProvider;
5use crate::types::split_namespaced_id;
6use crate::ToolError;
7
8#[derive(Debug, Clone, Copy, PartialEq, Eq)]
9pub enum LocalProviderName {
10    State,
11    Git,
12    #[cfg(feature = "openapi")]
13    Openapi,
14}
15
16#[derive(Debug, Clone, PartialEq)]
17pub struct LocalProviderCall {
18    pub provider: LocalProviderName,
19    pub method: String,
20    pub params: Value,
21}
22
23pub fn is_reserved_provider_namespace(namespace: &str) -> bool {
24    matches!(namespace, "state" | "git") || {
25        #[cfg(feature = "openapi")]
26        {
27            namespace == "openapi"
28        }
29        #[cfg(not(feature = "openapi"))]
30        {
31            let _ = namespace;
32            false
33        }
34    }
35}
36
37pub fn parse_local_provider_call(
38    id: &str,
39    params: Value,
40) -> Result<Option<LocalProviderCall>, ToolError> {
41    let Some((namespace, method)) = split_namespaced_id(id.trim()) else {
42        return Ok(None);
43    };
44    let provider = match namespace {
45        "state" => LocalProviderName::State,
46        "git" => LocalProviderName::Git,
47        #[cfg(feature = "openapi")]
48        "openapi" => LocalProviderName::Openapi,
49        _ => return Ok(None),
50    };
51    if method.trim().is_empty() {
52        return Err(ToolError::InvalidParam {
53            message: "local provider method must not be empty".to_string(),
54            param: "id".to_string(),
55        });
56    }
57    Ok(Some(LocalProviderCall {
58        provider,
59        method: method.to_string(),
60        params,
61    }))
62}
63
64pub async fn dispatch_local_provider(call: LocalProviderCall) -> Result<Value, ToolError> {
65    match call.provider {
66        LocalProviderName::State => {
67            StateProvider::default()
68                .dispatch(&call.method, call.params)
69                .await
70        }
71        LocalProviderName::Git => {
72            GitProvider::new(std::env::current_dir().map_err(|err| {
73                ToolError::internal_message(format!(
74                    "failed to resolve cwd for git provider: {err}"
75                ))
76            })?)
77            .dispatch(&call.method, call.params)
78            .await
79        }
80        #[cfg(feature = "openapi")]
81        LocalProviderName::Openapi => Err(ToolError::internal_message(
82            "openapi provider requires an OpenAPI registry and must be dispatched through openapi_feature",
83        )),
84    }
85}