break up acp/server.rs (#8932)
This commit is contained in:
+12
-1720
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,36 @@
|
||||
use super::*;
|
||||
|
||||
impl GooseAcpAgent {
|
||||
pub(super) async fn on_read_config(
|
||||
&self,
|
||||
req: ReadConfigRequest,
|
||||
) -> Result<ReadConfigResponse, sacp::Error> {
|
||||
let config = self.config()?;
|
||||
let response = match config.get_param::<serde_json::Value>(&req.key) {
|
||||
Ok(value) => ReadConfigResponse { value },
|
||||
Err(crate::config::ConfigError::NotFound(_)) => ReadConfigResponse {
|
||||
value: serde_json::Value::Null,
|
||||
},
|
||||
Err(e) => return Err(sacp::Error::internal_error().data(e.to_string())),
|
||||
};
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
pub(super) async fn on_upsert_config(
|
||||
&self,
|
||||
req: UpsertConfigRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
let config = self.config()?;
|
||||
config.set_param(&req.key, &req.value).internal_err()?;
|
||||
Ok(EmptyResponse {})
|
||||
}
|
||||
|
||||
pub(super) async fn on_remove_config(
|
||||
&self,
|
||||
req: RemoveConfigRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
let config = self.config()?;
|
||||
config.delete(&req.key).internal_err()?;
|
||||
Ok(EmptyResponse {})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,354 @@
|
||||
use super::*;
|
||||
use goose_acp_macros::custom_methods;
|
||||
|
||||
#[custom_methods]
|
||||
impl GooseAcpAgent {
|
||||
pub async fn dispatch_custom_request(
|
||||
&self,
|
||||
method: &str,
|
||||
params: serde_json::Value,
|
||||
) -> Result<serde_json::Value, sacp::Error> {
|
||||
self.handle_custom_request(method, params).await
|
||||
}
|
||||
|
||||
#[custom_method(AddExtensionRequest)]
|
||||
async fn dispatch_add_extension(
|
||||
&self,
|
||||
req: AddExtensionRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
self.on_add_extension(req).await
|
||||
}
|
||||
|
||||
#[custom_method(RemoveExtensionRequest)]
|
||||
async fn dispatch_remove_extension(
|
||||
&self,
|
||||
req: RemoveExtensionRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
self.on_remove_extension(req).await
|
||||
}
|
||||
|
||||
#[custom_method(GetToolsRequest)]
|
||||
async fn dispatch_get_tools(
|
||||
&self,
|
||||
req: GetToolsRequest,
|
||||
) -> Result<GetToolsResponse, sacp::Error> {
|
||||
self.on_get_tools(req).await
|
||||
}
|
||||
|
||||
#[custom_method(ReadResourceRequest)]
|
||||
async fn dispatch_read_resource(
|
||||
&self,
|
||||
req: ReadResourceRequest,
|
||||
) -> Result<ReadResourceResponse, sacp::Error> {
|
||||
self.on_read_resource(req).await
|
||||
}
|
||||
|
||||
#[custom_method(UpdateWorkingDirRequest)]
|
||||
async fn dispatch_update_working_dir(
|
||||
&self,
|
||||
req: UpdateWorkingDirRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
self.on_update_working_dir(req).await
|
||||
}
|
||||
|
||||
#[custom_method(DeleteSessionRequest)]
|
||||
async fn dispatch_delete_session(
|
||||
&self,
|
||||
req: DeleteSessionRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
self.on_delete_session(req).await
|
||||
}
|
||||
|
||||
#[custom_method(GetExtensionsRequest)]
|
||||
async fn dispatch_get_extensions(&self) -> Result<GetExtensionsResponse, sacp::Error> {
|
||||
self.on_get_extensions().await
|
||||
}
|
||||
|
||||
#[custom_method(AddConfigExtensionRequest)]
|
||||
async fn dispatch_add_config_extension(
|
||||
&self,
|
||||
req: AddConfigExtensionRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
self.on_add_config_extension(req).await
|
||||
}
|
||||
|
||||
#[custom_method(RemoveConfigExtensionRequest)]
|
||||
async fn dispatch_remove_config_extension(
|
||||
&self,
|
||||
req: RemoveConfigExtensionRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
self.on_remove_config_extension(req).await
|
||||
}
|
||||
|
||||
#[custom_method(ToggleConfigExtensionRequest)]
|
||||
async fn dispatch_toggle_config_extension(
|
||||
&self,
|
||||
req: ToggleConfigExtensionRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
self.on_toggle_config_extension(req).await
|
||||
}
|
||||
|
||||
#[custom_method(GetSessionExtensionsRequest)]
|
||||
async fn dispatch_get_session_extensions(
|
||||
&self,
|
||||
req: GetSessionExtensionsRequest,
|
||||
) -> Result<GetSessionExtensionsResponse, sacp::Error> {
|
||||
self.on_get_session_extensions(req).await
|
||||
}
|
||||
|
||||
#[custom_method(ListProvidersRequest)]
|
||||
async fn dispatch_list_providers(
|
||||
&self,
|
||||
req: ListProvidersRequest,
|
||||
) -> Result<ListProvidersResponse, sacp::Error> {
|
||||
self.on_list_providers(req).await
|
||||
}
|
||||
|
||||
#[custom_method(RefreshProviderInventoryRequest)]
|
||||
async fn dispatch_refresh_provider_inventory(
|
||||
&self,
|
||||
req: RefreshProviderInventoryRequest,
|
||||
) -> Result<RefreshProviderInventoryResponse, sacp::Error> {
|
||||
self.on_refresh_provider_inventory(req).await
|
||||
}
|
||||
|
||||
#[custom_method(ProviderConfigReadRequest)]
|
||||
async fn dispatch_read_provider_config(
|
||||
&self,
|
||||
req: ProviderConfigReadRequest,
|
||||
) -> Result<ProviderConfigReadResponse, sacp::Error> {
|
||||
self.on_read_provider_config(req).await
|
||||
}
|
||||
|
||||
#[custom_method(ProviderConfigStatusRequest)]
|
||||
async fn dispatch_provider_config_status(
|
||||
&self,
|
||||
req: ProviderConfigStatusRequest,
|
||||
) -> Result<ProviderConfigStatusResponse, sacp::Error> {
|
||||
self.on_provider_config_status(req).await
|
||||
}
|
||||
|
||||
#[custom_method(ProviderConfigSaveRequest)]
|
||||
async fn dispatch_save_provider_config(
|
||||
&self,
|
||||
req: ProviderConfigSaveRequest,
|
||||
) -> Result<ProviderConfigChangeResponse, sacp::Error> {
|
||||
self.on_save_provider_config(req).await
|
||||
}
|
||||
|
||||
#[custom_method(ProviderConfigDeleteRequest)]
|
||||
async fn dispatch_delete_provider_config(
|
||||
&self,
|
||||
req: ProviderConfigDeleteRequest,
|
||||
) -> Result<ProviderConfigChangeResponse, sacp::Error> {
|
||||
self.on_delete_provider_config(req).await
|
||||
}
|
||||
|
||||
#[custom_method(ReadConfigRequest)]
|
||||
async fn dispatch_read_config(
|
||||
&self,
|
||||
req: ReadConfigRequest,
|
||||
) -> Result<ReadConfigResponse, sacp::Error> {
|
||||
self.on_read_config(req).await
|
||||
}
|
||||
|
||||
#[custom_method(UpsertConfigRequest)]
|
||||
async fn dispatch_upsert_config(
|
||||
&self,
|
||||
req: UpsertConfigRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
self.on_upsert_config(req).await
|
||||
}
|
||||
|
||||
#[custom_method(RemoveConfigRequest)]
|
||||
async fn dispatch_remove_config(
|
||||
&self,
|
||||
req: RemoveConfigRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
self.on_remove_config(req).await
|
||||
}
|
||||
|
||||
#[custom_method(CheckSecretRequest)]
|
||||
async fn dispatch_check_secret(
|
||||
&self,
|
||||
req: CheckSecretRequest,
|
||||
) -> Result<CheckSecretResponse, sacp::Error> {
|
||||
self.on_check_secret(req).await
|
||||
}
|
||||
|
||||
#[custom_method(UpsertSecretRequest)]
|
||||
async fn dispatch_upsert_secret(
|
||||
&self,
|
||||
req: UpsertSecretRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
self.on_upsert_secret(req).await
|
||||
}
|
||||
|
||||
#[custom_method(RemoveSecretRequest)]
|
||||
async fn dispatch_remove_secret(
|
||||
&self,
|
||||
req: RemoveSecretRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
self.on_remove_secret(req).await
|
||||
}
|
||||
|
||||
#[custom_method(ExportSessionRequest)]
|
||||
async fn dispatch_export_session(
|
||||
&self,
|
||||
req: ExportSessionRequest,
|
||||
) -> Result<ExportSessionResponse, sacp::Error> {
|
||||
self.on_export_session(req).await
|
||||
}
|
||||
|
||||
#[custom_method(ImportSessionRequest)]
|
||||
async fn dispatch_import_session(
|
||||
&self,
|
||||
req: ImportSessionRequest,
|
||||
) -> Result<ImportSessionResponse, sacp::Error> {
|
||||
self.on_import_session(req).await
|
||||
}
|
||||
|
||||
#[custom_method(UpdateSessionProjectRequest)]
|
||||
async fn dispatch_update_session_project(
|
||||
&self,
|
||||
req: UpdateSessionProjectRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
self.on_update_session_project(req).await
|
||||
}
|
||||
|
||||
#[custom_method(RenameSessionRequest)]
|
||||
async fn dispatch_rename_session(
|
||||
&self,
|
||||
req: RenameSessionRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
self.on_rename_session(req).await
|
||||
}
|
||||
|
||||
#[custom_method(ArchiveSessionRequest)]
|
||||
async fn dispatch_archive_session(
|
||||
&self,
|
||||
req: ArchiveSessionRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
self.on_archive_session(req).await
|
||||
}
|
||||
|
||||
#[custom_method(UnarchiveSessionRequest)]
|
||||
async fn dispatch_unarchive_session(
|
||||
&self,
|
||||
req: UnarchiveSessionRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
self.on_unarchive_session(req).await
|
||||
}
|
||||
|
||||
#[custom_method(CreateSourceRequest)]
|
||||
async fn dispatch_create_source(
|
||||
&self,
|
||||
req: CreateSourceRequest,
|
||||
) -> Result<CreateSourceResponse, sacp::Error> {
|
||||
self.on_create_source(req).await
|
||||
}
|
||||
|
||||
#[custom_method(ListSourcesRequest)]
|
||||
async fn dispatch_list_sources(
|
||||
&self,
|
||||
req: ListSourcesRequest,
|
||||
) -> Result<ListSourcesResponse, sacp::Error> {
|
||||
self.on_list_sources(req).await
|
||||
}
|
||||
|
||||
#[custom_method(UpdateSourceRequest)]
|
||||
async fn dispatch_update_source(
|
||||
&self,
|
||||
req: UpdateSourceRequest,
|
||||
) -> Result<UpdateSourceResponse, sacp::Error> {
|
||||
self.on_update_source(req).await
|
||||
}
|
||||
|
||||
#[custom_method(DeleteSourceRequest)]
|
||||
async fn dispatch_delete_source(
|
||||
&self,
|
||||
req: DeleteSourceRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
self.on_delete_source(req).await
|
||||
}
|
||||
|
||||
#[custom_method(ExportSourceRequest)]
|
||||
async fn dispatch_export_source(
|
||||
&self,
|
||||
req: ExportSourceRequest,
|
||||
) -> Result<ExportSourceResponse, sacp::Error> {
|
||||
self.on_export_source(req).await
|
||||
}
|
||||
|
||||
#[custom_method(ImportSourcesRequest)]
|
||||
async fn dispatch_import_sources(
|
||||
&self,
|
||||
req: ImportSourcesRequest,
|
||||
) -> Result<ImportSourcesResponse, sacp::Error> {
|
||||
self.on_import_sources(req).await
|
||||
}
|
||||
|
||||
#[custom_method(DictationTranscribeRequest)]
|
||||
async fn dispatch_dictation_transcribe(
|
||||
&self,
|
||||
req: DictationTranscribeRequest,
|
||||
) -> Result<DictationTranscribeResponse, sacp::Error> {
|
||||
self.on_dictation_transcribe(req).await
|
||||
}
|
||||
|
||||
#[custom_method(DictationConfigRequest)]
|
||||
async fn dispatch_dictation_config(
|
||||
&self,
|
||||
_req: DictationConfigRequest,
|
||||
) -> Result<DictationConfigResponse, sacp::Error> {
|
||||
self.on_dictation_config(_req).await
|
||||
}
|
||||
|
||||
#[custom_method(DictationModelsListRequest)]
|
||||
async fn dispatch_dictation_models_list(
|
||||
&self,
|
||||
_req: DictationModelsListRequest,
|
||||
) -> Result<DictationModelsListResponse, sacp::Error> {
|
||||
self.on_dictation_models_list(_req).await
|
||||
}
|
||||
|
||||
#[custom_method(DictationModelDownloadRequest)]
|
||||
async fn dispatch_dictation_model_download(
|
||||
&self,
|
||||
_req: DictationModelDownloadRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
self.on_dictation_model_download(_req).await
|
||||
}
|
||||
|
||||
#[custom_method(DictationModelDownloadProgressRequest)]
|
||||
async fn dispatch_dictation_model_download_progress(
|
||||
&self,
|
||||
_req: DictationModelDownloadProgressRequest,
|
||||
) -> Result<DictationModelDownloadProgressResponse, sacp::Error> {
|
||||
self.on_dictation_model_download_progress(_req).await
|
||||
}
|
||||
|
||||
#[custom_method(DictationModelCancelRequest)]
|
||||
async fn dispatch_dictation_model_cancel(
|
||||
&self,
|
||||
_req: DictationModelCancelRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
self.on_dictation_model_cancel(_req).await
|
||||
}
|
||||
|
||||
#[custom_method(DictationModelDeleteRequest)]
|
||||
async fn dispatch_dictation_model_delete(
|
||||
&self,
|
||||
_req: DictationModelDeleteRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
self.on_dictation_model_delete(_req).await
|
||||
}
|
||||
|
||||
#[custom_method(DictationModelSelectRequest)]
|
||||
async fn dispatch_dictation_model_select(
|
||||
&self,
|
||||
req: DictationModelSelectRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
self.on_dictation_model_select(req).await
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,406 @@
|
||||
use super::*;
|
||||
#[cfg(feature = "local-inference")]
|
||||
use crate::dictation::providers::transcribe_local;
|
||||
use crate::dictation::providers::{
|
||||
all_providers, is_configured, transcribe_with_provider, DictationProvider,
|
||||
};
|
||||
#[cfg(feature = "local-inference")]
|
||||
use crate::dictation::whisper;
|
||||
|
||||
const OPENAI_TRANSCRIPTION_MODEL_CONFIG_KEY: &str = "OPENAI_TRANSCRIPTION_MODEL";
|
||||
const GROQ_TRANSCRIPTION_MODEL_CONFIG_KEY: &str = "GROQ_TRANSCRIPTION_MODEL";
|
||||
const ELEVENLABS_TRANSCRIPTION_MODEL_CONFIG_KEY: &str = "ELEVENLABS_TRANSCRIPTION_MODEL";
|
||||
const OPENAI_TRANSCRIPTION_MODEL: &str = "whisper-1";
|
||||
const GROQ_TRANSCRIPTION_MODEL: &str = "whisper-large-v3-turbo";
|
||||
const ELEVENLABS_TRANSCRIPTION_MODEL: &str = "scribe_v1";
|
||||
|
||||
impl GooseAcpAgent {
|
||||
pub(super) async fn on_dictation_transcribe(
|
||||
&self,
|
||||
req: DictationTranscribeRequest,
|
||||
) -> Result<DictationTranscribeResponse, sacp::Error> {
|
||||
use base64::{engine::general_purpose::STANDARD as BASE64, Engine};
|
||||
let config = crate::config::Config::global();
|
||||
|
||||
#[cfg(not(feature = "local-inference"))]
|
||||
if req.provider == "local" {
|
||||
return Err(sacp::Error::invalid_params()
|
||||
.data("Local inference is not available in this build"));
|
||||
}
|
||||
|
||||
let provider: DictationProvider = serde_json::from_value(serde_json::Value::String(
|
||||
req.provider.clone(),
|
||||
))
|
||||
.map_err(|_| {
|
||||
sacp::Error::invalid_params().data(format!("Unknown provider: {}", req.provider))
|
||||
})?;
|
||||
|
||||
let audio_bytes = BASE64
|
||||
.decode(&req.audio)
|
||||
.map_err(|_| sacp::Error::invalid_params().data("Invalid base64 audio data"))?;
|
||||
|
||||
if audio_bytes.len() > 50 * 1024 * 1024 {
|
||||
return Err(sacp::Error::invalid_params().data("Audio too large (max 50MB)"));
|
||||
}
|
||||
|
||||
let extension = match req.mime_type.as_str() {
|
||||
"audio/webm" | "audio/webm;codecs=opus" => "webm",
|
||||
"audio/mp4" => "mp4",
|
||||
"audio/mpeg" | "audio/mpga" => "mp3",
|
||||
"audio/m4a" => "m4a",
|
||||
"audio/wav" | "audio/x-wav" => "wav",
|
||||
other => {
|
||||
return Err(
|
||||
sacp::Error::invalid_params().data(format!("Unsupported format: {other}"))
|
||||
);
|
||||
}
|
||||
};
|
||||
|
||||
let text = match provider {
|
||||
#[cfg(feature = "local-inference")]
|
||||
DictationProvider::Local => transcribe_local(audio_bytes).await,
|
||||
remote => {
|
||||
let (model_param, default_model) = dictation_transcribe_params(remote);
|
||||
let model = dictation_selected_model(config, remote)
|
||||
.unwrap_or_else(|| default_model.to_string());
|
||||
transcribe_with_provider(
|
||||
remote,
|
||||
model_param.to_string(),
|
||||
model,
|
||||
audio_bytes,
|
||||
extension,
|
||||
&req.mime_type,
|
||||
)
|
||||
.await
|
||||
}
|
||||
}
|
||||
.internal_err()?;
|
||||
|
||||
Ok(DictationTranscribeResponse { text })
|
||||
}
|
||||
|
||||
pub(super) async fn on_dictation_config(
|
||||
&self,
|
||||
_req: DictationConfigRequest,
|
||||
) -> Result<DictationConfigResponse, sacp::Error> {
|
||||
let config = crate::config::Config::global();
|
||||
let mut providers = std::collections::HashMap::new();
|
||||
|
||||
for def in all_providers() {
|
||||
let provider = def.provider;
|
||||
let host = if let Some(host_key) = def.host_key {
|
||||
config
|
||||
.get(host_key, false)
|
||||
.ok()
|
||||
.and_then(|v| v.as_str().map(|s| s.to_string()))
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
let provider_key = serde_json::to_value(provider)
|
||||
.ok()
|
||||
.and_then(|v| v.as_str().map(|s| s.to_string()))
|
||||
.unwrap_or_else(|| format!("{:?}", provider).to_lowercase());
|
||||
providers.insert(
|
||||
provider_key,
|
||||
DictationProviderStatusEntry {
|
||||
configured: is_configured(provider),
|
||||
host,
|
||||
description: def.description.to_string(),
|
||||
uses_provider_config: def.uses_provider_config,
|
||||
settings_path: def.settings_path.map(|s| s.to_string()),
|
||||
config_key: if !def.uses_provider_config {
|
||||
Some(def.config_key.to_string())
|
||||
} else {
|
||||
None
|
||||
},
|
||||
model_config_key: dictation_model_config_key(provider),
|
||||
default_model: dictation_default_model(provider),
|
||||
selected_model: dictation_selected_model(config, provider),
|
||||
available_models: dictation_available_models(provider),
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
Ok(DictationConfigResponse { providers })
|
||||
}
|
||||
|
||||
pub(super) async fn on_dictation_models_list(
|
||||
&self,
|
||||
_req: DictationModelsListRequest,
|
||||
) -> Result<DictationModelsListResponse, sacp::Error> {
|
||||
#[cfg(feature = "local-inference")]
|
||||
{
|
||||
use crate::download_manager::{get_download_manager, DownloadStatus};
|
||||
|
||||
let manager = get_download_manager();
|
||||
let models = whisper::available_models()
|
||||
.iter()
|
||||
.map(|model| DictationLocalModelStatus {
|
||||
id: model.id.to_string(),
|
||||
label: model.id.to_string(),
|
||||
description: model.description.to_string(),
|
||||
size_mb: model.size_mb,
|
||||
downloaded: model.is_downloaded(),
|
||||
download_in_progress: manager
|
||||
.get_progress(model.id)
|
||||
.map(|progress| progress.status == DownloadStatus::Downloading)
|
||||
.unwrap_or(false),
|
||||
})
|
||||
.collect();
|
||||
|
||||
Ok(DictationModelsListResponse { models })
|
||||
}
|
||||
|
||||
#[cfg(not(feature = "local-inference"))]
|
||||
Ok(DictationModelsListResponse::default())
|
||||
}
|
||||
|
||||
pub(super) async fn on_dictation_model_download(
|
||||
&self,
|
||||
_req: DictationModelDownloadRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
#[cfg(feature = "local-inference")]
|
||||
{
|
||||
use crate::download_manager::get_download_manager;
|
||||
|
||||
let model = whisper::get_model(&_req.model_id)
|
||||
.ok_or_else(|| sacp::Error::invalid_params().data("Unknown model id"))?;
|
||||
let manager = get_download_manager();
|
||||
let model_id_for_config = model.id.to_string();
|
||||
|
||||
manager
|
||||
.download_model(
|
||||
model.id.to_string(),
|
||||
model.url.to_string(),
|
||||
model.local_path(),
|
||||
Some(Box::new(move || {
|
||||
let config = crate::config::Config::global();
|
||||
// Only auto-select this model if the user has no model
|
||||
// currently selected. This prevents silently switching
|
||||
// the active model mid-session when a user downloads an
|
||||
// additional model while one is already in use.
|
||||
let already_selected = config
|
||||
.get(whisper::LOCAL_WHISPER_MODEL_CONFIG_KEY, false)
|
||||
.ok()
|
||||
.and_then(|value| value.as_str().map(str::to_owned))
|
||||
.filter(|model_id| {
|
||||
// Treat a deleted model file as no active selection
|
||||
// so a fresh download can auto-select cleanly.
|
||||
whisper::get_model(model_id)
|
||||
.is_some_and(|model| model.is_downloaded())
|
||||
});
|
||||
if already_selected.is_none() {
|
||||
if let Err(e) = config.set_param(
|
||||
whisper::LOCAL_WHISPER_MODEL_CONFIG_KEY,
|
||||
model_id_for_config.clone(),
|
||||
) {
|
||||
error!("Failed to save LOCAL_WHISPER_MODEL after download: {}", e);
|
||||
}
|
||||
}
|
||||
})),
|
||||
)
|
||||
.await
|
||||
.internal_err()?;
|
||||
|
||||
Ok(EmptyResponse {})
|
||||
}
|
||||
|
||||
#[cfg(not(feature = "local-inference"))]
|
||||
Err(sacp::Error::invalid_params().data("Local inference not enabled"))
|
||||
}
|
||||
|
||||
pub(super) async fn on_dictation_model_download_progress(
|
||||
&self,
|
||||
_req: DictationModelDownloadProgressRequest,
|
||||
) -> Result<DictationModelDownloadProgressResponse, sacp::Error> {
|
||||
#[cfg(feature = "local-inference")]
|
||||
{
|
||||
use crate::download_manager::get_download_manager;
|
||||
|
||||
let manager = get_download_manager();
|
||||
let progress =
|
||||
manager
|
||||
.get_progress(&_req.model_id)
|
||||
.map(|progress| DictationDownloadProgress {
|
||||
bytes_downloaded: progress.bytes_downloaded,
|
||||
total_bytes: progress.total_bytes,
|
||||
progress_percent: progress.progress_percent,
|
||||
status: serde_json::to_value(&progress.status)
|
||||
.ok()
|
||||
.and_then(|value| value.as_str().map(ToOwned::to_owned))
|
||||
.unwrap_or_else(|| "unknown".to_string()),
|
||||
error: progress.error,
|
||||
});
|
||||
|
||||
Ok(DictationModelDownloadProgressResponse { progress })
|
||||
}
|
||||
|
||||
#[cfg(not(feature = "local-inference"))]
|
||||
Ok(DictationModelDownloadProgressResponse { progress: None })
|
||||
}
|
||||
|
||||
pub(super) async fn on_dictation_model_cancel(
|
||||
&self,
|
||||
_req: DictationModelCancelRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
#[cfg(feature = "local-inference")]
|
||||
{
|
||||
use crate::download_manager::get_download_manager;
|
||||
|
||||
let manager = get_download_manager();
|
||||
manager.cancel_download(&_req.model_id).internal_err()?;
|
||||
|
||||
Ok(EmptyResponse {})
|
||||
}
|
||||
|
||||
#[cfg(not(feature = "local-inference"))]
|
||||
Err(sacp::Error::invalid_params().data("Local inference not enabled"))
|
||||
}
|
||||
|
||||
pub(super) async fn on_dictation_model_delete(
|
||||
&self,
|
||||
_req: DictationModelDeleteRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
#[cfg(feature = "local-inference")]
|
||||
{
|
||||
let model = whisper::get_model(&_req.model_id)
|
||||
.ok_or_else(|| sacp::Error::invalid_params().data("Unknown model id"))?;
|
||||
let path = model.local_path();
|
||||
|
||||
if !path.exists() {
|
||||
return Err(sacp::Error::invalid_params().data("Model not downloaded"));
|
||||
}
|
||||
|
||||
std::fs::remove_file(path).internal_err()?;
|
||||
|
||||
Ok(EmptyResponse {})
|
||||
}
|
||||
|
||||
#[cfg(not(feature = "local-inference"))]
|
||||
Err(sacp::Error::invalid_params().data("Local inference not enabled"))
|
||||
}
|
||||
|
||||
pub(super) async fn on_dictation_model_select(
|
||||
&self,
|
||||
req: DictationModelSelectRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
#[cfg(not(feature = "local-inference"))]
|
||||
if req.provider == "local" {
|
||||
return Err(sacp::Error::invalid_params().data("Local inference not enabled"));
|
||||
}
|
||||
|
||||
let provider: DictationProvider = serde_json::from_value(serde_json::Value::String(
|
||||
req.provider.clone(),
|
||||
))
|
||||
.map_err(|_| {
|
||||
sacp::Error::invalid_params().data(format!("Unknown provider: {}", req.provider))
|
||||
})?;
|
||||
|
||||
let key = match provider {
|
||||
DictationProvider::OpenAI => OPENAI_TRANSCRIPTION_MODEL_CONFIG_KEY,
|
||||
DictationProvider::Groq => GROQ_TRANSCRIPTION_MODEL_CONFIG_KEY,
|
||||
DictationProvider::ElevenLabs => ELEVENLABS_TRANSCRIPTION_MODEL_CONFIG_KEY,
|
||||
#[cfg(feature = "local-inference")]
|
||||
DictationProvider::Local => {
|
||||
let model = whisper::get_model(&req.model_id)
|
||||
.ok_or_else(|| sacp::Error::invalid_params().data("Unknown model id"))?;
|
||||
if !model.is_downloaded() {
|
||||
return Err(
|
||||
sacp::Error::invalid_params().data("Local Whisper model is not downloaded")
|
||||
);
|
||||
}
|
||||
whisper::LOCAL_WHISPER_MODEL_CONFIG_KEY
|
||||
}
|
||||
};
|
||||
|
||||
crate::config::Config::global()
|
||||
.set_param(key, req.model_id)
|
||||
.internal_err()?;
|
||||
|
||||
Ok(EmptyResponse {})
|
||||
}
|
||||
}
|
||||
fn dictation_model_config_key(provider: DictationProvider) -> Option<String> {
|
||||
match provider {
|
||||
DictationProvider::OpenAI => Some(OPENAI_TRANSCRIPTION_MODEL_CONFIG_KEY.to_string()),
|
||||
DictationProvider::Groq => Some(GROQ_TRANSCRIPTION_MODEL_CONFIG_KEY.to_string()),
|
||||
DictationProvider::ElevenLabs => {
|
||||
Some(ELEVENLABS_TRANSCRIPTION_MODEL_CONFIG_KEY.to_string())
|
||||
}
|
||||
#[cfg(feature = "local-inference")]
|
||||
DictationProvider::Local => Some(whisper::LOCAL_WHISPER_MODEL_CONFIG_KEY.to_string()),
|
||||
}
|
||||
}
|
||||
|
||||
/// Returns the (param_name, default_model) pair used by `transcribe_with_provider`
|
||||
/// for remote dictation providers. Local inference is handled separately.
|
||||
fn dictation_transcribe_params(provider: DictationProvider) -> (&'static str, &'static str) {
|
||||
match provider {
|
||||
DictationProvider::OpenAI => ("model", OPENAI_TRANSCRIPTION_MODEL),
|
||||
DictationProvider::Groq => ("model", GROQ_TRANSCRIPTION_MODEL),
|
||||
DictationProvider::ElevenLabs => ("model_id", ELEVENLABS_TRANSCRIPTION_MODEL),
|
||||
#[cfg(feature = "local-inference")]
|
||||
DictationProvider::Local => ("", ""),
|
||||
}
|
||||
}
|
||||
|
||||
fn dictation_default_model(provider: DictationProvider) -> Option<String> {
|
||||
match provider {
|
||||
DictationProvider::OpenAI => Some(OPENAI_TRANSCRIPTION_MODEL.to_string()),
|
||||
DictationProvider::Groq => Some(GROQ_TRANSCRIPTION_MODEL.to_string()),
|
||||
DictationProvider::ElevenLabs => Some(ELEVENLABS_TRANSCRIPTION_MODEL.to_string()),
|
||||
#[cfg(feature = "local-inference")]
|
||||
DictationProvider::Local => Some(whisper::recommend_model().to_string()),
|
||||
}
|
||||
}
|
||||
|
||||
fn dictation_selected_model(config: &Config, provider: DictationProvider) -> Option<String> {
|
||||
#[cfg(feature = "local-inference")]
|
||||
if provider == DictationProvider::Local {
|
||||
return config
|
||||
.get(whisper::LOCAL_WHISPER_MODEL_CONFIG_KEY, false)
|
||||
.ok()
|
||||
.and_then(|value| value.as_str().map(str::to_owned))
|
||||
.filter(|model_id| whisper::get_model(model_id).is_some())
|
||||
.or_else(|| dictation_default_model(provider));
|
||||
}
|
||||
|
||||
dictation_model_config_key(provider)
|
||||
.and_then(|key| {
|
||||
config
|
||||
.get(&key, false)
|
||||
.ok()
|
||||
.and_then(|value| value.as_str().map(str::to_owned))
|
||||
})
|
||||
.or_else(|| dictation_default_model(provider))
|
||||
}
|
||||
|
||||
fn dictation_available_models(provider: DictationProvider) -> Vec<DictationModelOption> {
|
||||
match provider {
|
||||
DictationProvider::OpenAI => vec![DictationModelOption {
|
||||
id: OPENAI_TRANSCRIPTION_MODEL.to_string(),
|
||||
label: "Whisper-1".to_string(),
|
||||
description: "OpenAI's hosted Whisper transcription model.".to_string(),
|
||||
}],
|
||||
DictationProvider::Groq => vec![DictationModelOption {
|
||||
id: GROQ_TRANSCRIPTION_MODEL.to_string(),
|
||||
label: "Whisper Large V3 Turbo".to_string(),
|
||||
description: "Groq's fast hosted Whisper transcription model.".to_string(),
|
||||
}],
|
||||
DictationProvider::ElevenLabs => vec![DictationModelOption {
|
||||
id: ELEVENLABS_TRANSCRIPTION_MODEL.to_string(),
|
||||
label: "Scribe v1".to_string(),
|
||||
description: "ElevenLabs' hosted speech-to-text model.".to_string(),
|
||||
}],
|
||||
#[cfg(feature = "local-inference")]
|
||||
DictationProvider::Local => whisper::available_models()
|
||||
.iter()
|
||||
.map(|model| DictationModelOption {
|
||||
id: model.id.to_string(),
|
||||
label: model.id.to_string(),
|
||||
description: model.description.to_string(),
|
||||
})
|
||||
.collect(),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,397 @@
|
||||
use super::*;
|
||||
|
||||
impl HandleDispatchFrom<Client> for GooseAcpHandler {
|
||||
fn describe_chain(&self) -> impl std::fmt::Debug {
|
||||
"goose-acp"
|
||||
}
|
||||
|
||||
fn handle_dispatch_from(
|
||||
&mut self,
|
||||
message: Dispatch,
|
||||
cx: ConnectionTo<Client>,
|
||||
) -> impl std::future::Future<Output = Result<Handled<Dispatch>, sacp::Error>> + Send {
|
||||
let agent = self.agent.clone();
|
||||
|
||||
// The MatchDispatchFrom chain produces an ~85KB async state machine.
|
||||
// Box::pin moves it to the heap so it doesn't overflow the tokio worker stack.
|
||||
Box::pin(async move {
|
||||
// InitializeRequest runs inline: it sets connection-scoped state
|
||||
// (client fs/terminal capabilities) that later handlers read with
|
||||
// defaults, so a pipelined NewSessionRequest must not race ahead of it.
|
||||
MatchDispatchFrom::new(message, &cx)
|
||||
.if_request(
|
||||
|req: InitializeRequest, responder: Responder<InitializeResponse>| async {
|
||||
responder.respond_with_result(agent.on_initialize(req).await)
|
||||
},
|
||||
)
|
||||
.await
|
||||
.if_request(
|
||||
|_req: AuthenticateRequest, responder: Responder<AuthenticateResponse>| async {
|
||||
responder.respond(AuthenticateResponse::new())
|
||||
},
|
||||
)
|
||||
.await
|
||||
.if_request(
|
||||
|req: NewSessionRequest, responder: Responder<NewSessionResponse>| async {
|
||||
let agent = agent.clone();
|
||||
let cx_clone = cx.clone();
|
||||
cx.spawn(async move {
|
||||
responder.respond_with_result(agent.on_new_session(&cx_clone, req).await)?;
|
||||
Ok(())
|
||||
})?;
|
||||
Ok(())
|
||||
},
|
||||
)
|
||||
.await
|
||||
.if_request(
|
||||
|req: LoadSessionRequest, responder: Responder<LoadSessionResponse>| async {
|
||||
let agent = agent.clone();
|
||||
let cx_clone = cx.clone();
|
||||
cx.spawn(async move {
|
||||
match agent.on_load_session(&cx_clone, req).await {
|
||||
Ok(response) => {
|
||||
responder.respond(response)?;
|
||||
}
|
||||
Err(e) => {
|
||||
responder.respond_with_error(e)?;
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
})?;
|
||||
Ok(())
|
||||
},
|
||||
)
|
||||
.await
|
||||
.if_request(
|
||||
|req: PromptRequest, responder: Responder<PromptResponse>| async {
|
||||
let agent = agent.clone();
|
||||
let cx_clone = cx.clone();
|
||||
cx.spawn(async move {
|
||||
match agent.on_prompt(&cx_clone, req).await {
|
||||
Ok(response) => {
|
||||
responder.respond(response)?;
|
||||
}
|
||||
Err(e) => {
|
||||
responder.respond_with_error(e)?;
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
})?;
|
||||
Ok(())
|
||||
},
|
||||
)
|
||||
.await
|
||||
.if_notification(|notif: CancelNotification| async {
|
||||
let agent = agent.clone();
|
||||
agent.on_cancel(notif).await?;
|
||||
Ok(())
|
||||
})
|
||||
.await
|
||||
// set_config_option (SACP 11) and legacy set_mode/set_model; custom _goose/* in otherwise.
|
||||
.if_request({
|
||||
let agent = agent.clone();
|
||||
let cx = cx.clone();
|
||||
|req: SetSessionConfigOptionRequest, responder: Responder<SetSessionConfigOptionResponse>| async move {
|
||||
let cx_spawn = cx.clone();
|
||||
cx.spawn(async move {
|
||||
let cx = cx_spawn;
|
||||
let value_id = req.value.as_value_id()
|
||||
.ok_or_else(|| sacp::Error::invalid_params().data("Expected a value ID"))?
|
||||
.clone();
|
||||
let session_id = req.session_id.clone();
|
||||
let sid = sid_short(session_id.0.as_ref());
|
||||
let config_id = req.config_id.0.to_string();
|
||||
let t_handler = std::time::Instant::now();
|
||||
match config_id.as_ref() {
|
||||
"provider" => {
|
||||
Config::global().invalidate_secrets_cache();
|
||||
match agent.update_provider(&session_id.0, &value_id.0, None, None, None).await {
|
||||
Ok(_) => {}
|
||||
Err(e) => { responder.respond_with_error(e)?; return Ok(()); }
|
||||
}
|
||||
}
|
||||
"mode" => {
|
||||
match agent.on_set_mode(&session_id.0, &value_id.0).await {
|
||||
Ok(_) => {}
|
||||
Err(e) => { responder.respond_with_error(e)?; return Ok(()); }
|
||||
}
|
||||
}
|
||||
"model" => {
|
||||
match agent.on_set_model(&session_id.0, &value_id.0).await {
|
||||
Ok(_) => {}
|
||||
Err(e) => { responder.respond_with_error(e)?; return Ok(()); }
|
||||
}
|
||||
}
|
||||
other => {
|
||||
responder.respond_with_error(
|
||||
sacp::Error::invalid_params().data(format!("Unsupported config option: {}", other))
|
||||
)?;
|
||||
return Ok(());
|
||||
}
|
||||
}
|
||||
// Respond immediately using the current provider inventory snapshot.
|
||||
let (notification, config_options) = agent.build_config_update(&session_id).await?;
|
||||
cx.send_notification(notification)?;
|
||||
responder.respond(SetSessionConfigOptionResponse::new(config_options))?;
|
||||
|
||||
let maybe_refresh = if config_id == "provider" {
|
||||
let provider_id = value_id.0.to_string();
|
||||
agent
|
||||
.provider_inventory
|
||||
.plan_refresh_jobs(std::slice::from_ref(&provider_id))
|
||||
.await
|
||||
.ok()
|
||||
.and_then(|plan| {
|
||||
plan.started
|
||||
.into_iter()
|
||||
.find(|job| job.provider_id == provider_id)
|
||||
})
|
||||
} else {
|
||||
None
|
||||
};
|
||||
if let Some(refresh_job) = maybe_refresh {
|
||||
let agent_bg = agent.clone();
|
||||
let cx_bg = cx.clone();
|
||||
let session_id_bg = session_id.clone();
|
||||
tokio::spawn(async move {
|
||||
let refresh_identity = refresh_job.identity;
|
||||
let refresh_provider_id = refresh_job.provider_id;
|
||||
let mut refresh_guard =
|
||||
agent_bg.provider_inventory.refresh_guard(&refresh_identity);
|
||||
let provider_result: Result<Arc<dyn Provider>> =
|
||||
AssertUnwindSafe(async {
|
||||
let session_agent =
|
||||
agent_bg.get_session_agent(&session_id_bg.0, None).await?;
|
||||
let provider = session_agent
|
||||
.provider()
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!(e.to_string()))?;
|
||||
let provider_name = provider.get_name().to_string();
|
||||
if provider_name != refresh_provider_id {
|
||||
return Err(anyhow::anyhow!(
|
||||
"provider changed before inventory refresh completed"
|
||||
));
|
||||
}
|
||||
Ok(provider)
|
||||
})
|
||||
.catch_unwind()
|
||||
.await
|
||||
.map_err(|_| {
|
||||
anyhow::anyhow!("provider inventory refresh task panicked")
|
||||
})
|
||||
.and_then(|result| result);
|
||||
|
||||
let fetch_result = match provider_result {
|
||||
Ok(provider) => {
|
||||
match ensure_refresh_identity_current(
|
||||
&refresh_provider_id,
|
||||
&refresh_identity,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(()) => match AssertUnwindSafe(
|
||||
provider.fetch_recommended_models(),
|
||||
)
|
||||
.catch_unwind()
|
||||
.await
|
||||
{
|
||||
Ok(Ok(models)) => Ok(models),
|
||||
Ok(Err(error)) => {
|
||||
Err(anyhow::anyhow!(error.to_string()))
|
||||
}
|
||||
Err(_) => Err(anyhow::anyhow!(
|
||||
"provider inventory refresh task panicked"
|
||||
)),
|
||||
},
|
||||
Err(error) => Err(error),
|
||||
}
|
||||
}
|
||||
Err(error) => Err(error),
|
||||
};
|
||||
|
||||
match fetch_result {
|
||||
Ok(models) => match agent_bg
|
||||
.provider_inventory
|
||||
.store_refreshed_models_for_identity(
|
||||
&refresh_identity,
|
||||
&models,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(()) => {
|
||||
refresh_guard.complete();
|
||||
match agent_bg.build_config_update(&session_id_bg).await
|
||||
{
|
||||
Ok((fresh_notification, _)) => {
|
||||
let _ = cx_bg
|
||||
.send_notification(fresh_notification);
|
||||
}
|
||||
Err(error) => warn!(
|
||||
provider = %refresh_provider_id,
|
||||
error = %error,
|
||||
"failed to build config update after provider inventory refresh"
|
||||
),
|
||||
}
|
||||
}
|
||||
Err(error) => warn!(
|
||||
provider = %refresh_provider_id,
|
||||
error = %error,
|
||||
"failed to store refreshed provider inventory after config change"
|
||||
),
|
||||
},
|
||||
Err(error) => {
|
||||
let error_message = error.to_string();
|
||||
match agent_bg
|
||||
.provider_inventory
|
||||
.store_refresh_error_for_identity(
|
||||
&refresh_identity,
|
||||
error_message.clone(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(()) => refresh_guard.complete(),
|
||||
Err(store_error) => warn!(
|
||||
provider = %refresh_provider_id,
|
||||
error = %store_error,
|
||||
refresh_error = %error_message,
|
||||
"failed to store provider inventory refresh error after config change"
|
||||
),
|
||||
}
|
||||
warn!(
|
||||
provider = %refresh_provider_id,
|
||||
error = %error_message,
|
||||
"provider inventory refresh failed after config change"
|
||||
);
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
debug!(target: "perf", sid = %sid, ms = t_handler.elapsed().as_millis() as u64, config_id = %config_id, "perf: set_config_option done");
|
||||
Ok(())
|
||||
})?;
|
||||
Ok(())
|
||||
}
|
||||
})
|
||||
.await
|
||||
.if_request({
|
||||
let agent = agent.clone();
|
||||
let cx = cx.clone();
|
||||
|req: SetSessionModeRequest, responder: Responder<SetSessionModeResponse>| async move {
|
||||
let cx_spawn = cx.clone();
|
||||
cx.spawn(async move {
|
||||
let cx = cx_spawn;
|
||||
let session_id = req.session_id.clone();
|
||||
let mode_id = req.mode_id.clone();
|
||||
match agent.on_set_mode(&session_id.0, &mode_id.0).await {
|
||||
Ok(resp) => {
|
||||
// Notify before responding so clients see the mode update before block_task unblocks.
|
||||
cx.send_notification(SessionNotification::new(
|
||||
session_id,
|
||||
SessionUpdate::CurrentModeUpdate(
|
||||
CurrentModeUpdate::new(mode_id),
|
||||
),
|
||||
))?;
|
||||
responder.respond(resp)?;
|
||||
}
|
||||
Err(e) => {
|
||||
responder.respond_with_error(e)?;
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
})?;
|
||||
Ok(())
|
||||
}
|
||||
})
|
||||
.await
|
||||
.if_request({
|
||||
let agent = agent.clone();
|
||||
let cx = cx.clone();
|
||||
|req: SetSessionModelRequest, responder: Responder<SetSessionModelResponse>| async move {
|
||||
let cx_spawn = cx.clone();
|
||||
cx.spawn(async move {
|
||||
let cx = cx_spawn;
|
||||
let session_id = req.session_id.clone();
|
||||
match agent.on_set_model(&session_id.0, &req.model_id.0).await {
|
||||
Ok(resp) => {
|
||||
let (notification, _) = agent.build_config_update(&session_id).await?;
|
||||
cx.send_notification(notification)?;
|
||||
responder.respond(resp)?;
|
||||
}
|
||||
Err(e) => responder.respond_with_error(e)?,
|
||||
}
|
||||
Ok(())
|
||||
})?;
|
||||
Ok(())
|
||||
}
|
||||
})
|
||||
.await
|
||||
.if_request({
|
||||
let agent = agent.clone();
|
||||
let cx = cx.clone();
|
||||
|_req: ListSessionsRequest, responder: Responder<ListSessionsResponse>| async move {
|
||||
cx.spawn(async move {
|
||||
responder.respond(agent.on_list_sessions().await?)?;
|
||||
Ok(())
|
||||
})?;
|
||||
Ok(())
|
||||
}
|
||||
})
|
||||
.await
|
||||
.if_request({
|
||||
let agent = agent.clone();
|
||||
let cx = cx.clone();
|
||||
|req: CloseSessionRequest, responder: Responder<CloseSessionResponse>| async move {
|
||||
cx.spawn(async move {
|
||||
responder.respond(agent.on_close_session(&req.session_id.0).await?)?;
|
||||
Ok(())
|
||||
})?;
|
||||
Ok(())
|
||||
}
|
||||
})
|
||||
.await
|
||||
.if_request({
|
||||
let agent = agent.clone();
|
||||
let cx = cx.clone();
|
||||
|req: ForkSessionRequest, responder: Responder<ForkSessionResponse>| async move {
|
||||
let cx_spawn = cx.clone();
|
||||
cx.spawn(async move {
|
||||
responder.respond_with_result(agent.on_fork_session(&cx_spawn, req).await)?;
|
||||
Ok(())
|
||||
})?;
|
||||
Ok(())
|
||||
}
|
||||
})
|
||||
.await
|
||||
.otherwise({
|
||||
let agent = agent.clone();
|
||||
let cx = cx.clone();
|
||||
|message: Dispatch| async move {
|
||||
match message {
|
||||
Dispatch::Request(req, responder) => {
|
||||
cx.spawn(async move {
|
||||
match agent.dispatch_custom_request(&req.method, req.params).await {
|
||||
Ok(json) => responder.respond(json)?,
|
||||
Err(e) => responder.respond_with_error(e)?,
|
||||
}
|
||||
Ok(())
|
||||
})?;
|
||||
Ok(())
|
||||
}
|
||||
Dispatch::Response(result, router) => {
|
||||
debug!(method = %router.method(), id = %router.id(), ok = result.is_ok(), "routing response");
|
||||
router.respond_with_result(result)?;
|
||||
Ok(())
|
||||
}
|
||||
Dispatch::Notification(notif) => {
|
||||
debug!(method = %notif.method, "unhandled notification");
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
.await
|
||||
.map(|()| Handled::Yes)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,136 @@
|
||||
use super::*;
|
||||
|
||||
impl GooseAcpAgent {
|
||||
pub(super) async fn on_add_extension(
|
||||
&self,
|
||||
req: AddExtensionRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
let internal_id = self.internal_session_id(&req.session_id).await?;
|
||||
let config: ExtensionConfig = serde_json::from_value(req.config)
|
||||
.map_err(|e| sacp::Error::invalid_params().data(format!("bad config: {e}")))?;
|
||||
let agent = self.get_session_agent(&req.session_id, None).await?;
|
||||
agent
|
||||
.add_extension(config, &internal_id)
|
||||
.await
|
||||
.internal_err()?;
|
||||
Ok(EmptyResponse {})
|
||||
}
|
||||
|
||||
pub(super) async fn on_remove_extension(
|
||||
&self,
|
||||
req: RemoveExtensionRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
let internal_id = self.internal_session_id(&req.session_id).await?;
|
||||
let agent = self.get_session_agent(&req.session_id, None).await?;
|
||||
agent
|
||||
.remove_extension(&req.name, &internal_id)
|
||||
.await
|
||||
.internal_err()?;
|
||||
Ok(EmptyResponse {})
|
||||
}
|
||||
|
||||
pub(super) async fn on_get_extensions(&self) -> Result<GetExtensionsResponse, sacp::Error> {
|
||||
let extensions = crate::config::extensions::get_all_extensions();
|
||||
let warnings = crate::config::extensions::get_warnings();
|
||||
let extensions_json = extensions
|
||||
.into_iter()
|
||||
.map(|e| {
|
||||
let config_key = e.config.key();
|
||||
let mut value = serde_json::to_value(&e)?;
|
||||
if let Some(obj) = value.as_object_mut() {
|
||||
obj.insert(
|
||||
"config_key".to_string(),
|
||||
serde_json::Value::String(config_key),
|
||||
);
|
||||
}
|
||||
Ok::<_, serde_json::Error>(value)
|
||||
})
|
||||
.collect::<Result<Vec<_>, _>>()
|
||||
.internal_err()?;
|
||||
Ok(GetExtensionsResponse {
|
||||
extensions: extensions_json,
|
||||
warnings,
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) async fn on_add_config_extension(
|
||||
&self,
|
||||
req: AddConfigExtensionRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
let mut obj = match req.extension_config {
|
||||
serde_json::Value::Object(obj) => obj,
|
||||
_ => {
|
||||
return Err(
|
||||
sacp::Error::invalid_params().data("extensionConfig must be a JSON object")
|
||||
);
|
||||
}
|
||||
};
|
||||
obj.insert(
|
||||
"name".to_string(),
|
||||
serde_json::Value::String(req.name.clone()),
|
||||
);
|
||||
|
||||
let config: crate::agents::ExtensionConfig =
|
||||
serde_json::from_value(serde_json::Value::Object(obj))
|
||||
.map_err(|e| sacp::Error::invalid_params().data(format!("bad config: {e}")))?;
|
||||
|
||||
crate::config::extensions::set_extension(crate::config::extensions::ExtensionEntry {
|
||||
enabled: req.enabled,
|
||||
config,
|
||||
});
|
||||
Ok(EmptyResponse {})
|
||||
}
|
||||
|
||||
pub(super) async fn on_remove_config_extension(
|
||||
&self,
|
||||
req: RemoveConfigExtensionRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
let keys = crate::config::extensions::get_all_extension_names();
|
||||
if !keys.iter().any(|k| k == &req.config_key) {
|
||||
return Err(sacp::Error::invalid_params()
|
||||
.data(format!("Extension '{}' not found", req.config_key)));
|
||||
}
|
||||
crate::config::extensions::remove_extension(&req.config_key);
|
||||
Ok(EmptyResponse {})
|
||||
}
|
||||
|
||||
pub(super) async fn on_toggle_config_extension(
|
||||
&self,
|
||||
req: ToggleConfigExtensionRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
let keys = crate::config::extensions::get_all_extension_names();
|
||||
if !keys.iter().any(|k| k == &req.config_key) {
|
||||
return Err(sacp::Error::invalid_params()
|
||||
.data(format!("Extension '{}' not found", req.config_key)));
|
||||
}
|
||||
crate::config::extensions::set_extension_enabled(&req.config_key, req.enabled);
|
||||
Ok(EmptyResponse {})
|
||||
}
|
||||
|
||||
pub(super) async fn on_get_session_extensions(
|
||||
&self,
|
||||
req: GetSessionExtensionsRequest,
|
||||
) -> Result<GetSessionExtensionsResponse, sacp::Error> {
|
||||
let internal_id = self.internal_session_id(&req.session_id).await?;
|
||||
let session = self
|
||||
.session_manager
|
||||
.get_session(&internal_id, false)
|
||||
.await
|
||||
.internal_err()?;
|
||||
|
||||
let extensions = EnabledExtensionsState::extensions_or_default(
|
||||
Some(&session.extension_data),
|
||||
crate::config::Config::global(),
|
||||
);
|
||||
|
||||
let extensions_json = extensions
|
||||
.into_iter()
|
||||
.map(|e| serde_json::to_value(&e))
|
||||
.collect::<Result<Vec<_>, _>>()
|
||||
.internal_err()?;
|
||||
|
||||
Ok(GetSessionExtensionsResponse {
|
||||
extensions: extensions_json,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,420 @@
|
||||
use super::*;
|
||||
|
||||
fn inventory_entry_to_dto(entry: ProviderInventoryEntry) -> ProviderInventoryEntryDto {
|
||||
let stale = ProviderInventoryService::is_stale(&entry);
|
||||
ProviderInventoryEntryDto {
|
||||
provider_id: entry.provider_id,
|
||||
provider_name: entry.provider_name,
|
||||
description: entry.description,
|
||||
default_model: entry.default_model,
|
||||
configured: entry.configured,
|
||||
provider_type: format!("{:?}", entry.provider_type),
|
||||
config_keys: entry
|
||||
.config_keys
|
||||
.into_iter()
|
||||
.map(provider_config_key_to_dto)
|
||||
.collect(),
|
||||
setup_steps: entry.setup_steps,
|
||||
supports_refresh: entry.supports_refresh,
|
||||
refreshing: entry.refreshing,
|
||||
models: entry
|
||||
.models
|
||||
.into_iter()
|
||||
.map(|m| ProviderInventoryModelDto {
|
||||
id: m.id,
|
||||
name: m.name,
|
||||
family: m.family,
|
||||
context_limit: m.context_limit,
|
||||
reasoning: m.reasoning,
|
||||
recommended: m.recommended,
|
||||
})
|
||||
.collect(),
|
||||
last_updated_at: entry.last_updated_at.map(|t| t.to_rfc3339()),
|
||||
last_refresh_attempt_at: entry.last_refresh_attempt_at.map(|t| t.to_rfc3339()),
|
||||
last_refresh_error: entry.last_refresh_error,
|
||||
stale,
|
||||
model_selection_hint: entry.model_selection_hint,
|
||||
}
|
||||
}
|
||||
|
||||
fn provider_config_key_to_dto(key: crate::providers::base::ConfigKey) -> ProviderConfigKey {
|
||||
ProviderConfigKey {
|
||||
name: key.name,
|
||||
required: key.required,
|
||||
secret: key.secret,
|
||||
default: key.default,
|
||||
oauth_flow: key.oauth_flow,
|
||||
device_code_flow: key.device_code_flow,
|
||||
primary: key.primary,
|
||||
}
|
||||
}
|
||||
|
||||
const SECRET_MASK_PREFIX_LEN: usize = 4;
|
||||
const SECRET_MASK_SUFFIX_LEN: usize = 3;
|
||||
const SECRET_MASK_FALLBACK: &str = "***";
|
||||
|
||||
fn mask_secret_value(value: &str) -> String {
|
||||
let prefix: String = value.chars().take(SECRET_MASK_PREFIX_LEN).collect();
|
||||
let suffix_chars: Vec<char> = value.chars().rev().take(SECRET_MASK_SUFFIX_LEN).collect();
|
||||
let suffix: String = suffix_chars.into_iter().rev().collect();
|
||||
|
||||
if prefix.is_empty()
|
||||
|| suffix.is_empty()
|
||||
|| value.chars().count() <= SECRET_MASK_PREFIX_LEN + SECRET_MASK_SUFFIX_LEN
|
||||
{
|
||||
return SECRET_MASK_FALLBACK.to_string();
|
||||
}
|
||||
|
||||
format!("{prefix}...{suffix}")
|
||||
}
|
||||
|
||||
fn config_value_to_string(value: &serde_json::Value) -> Option<String> {
|
||||
match value {
|
||||
serde_json::Value::Null => None,
|
||||
serde_json::Value::String(value) if value.is_empty() => None,
|
||||
serde_json::Value::String(value) => Some(value.clone()),
|
||||
other => serde_json::to_string(other).ok(),
|
||||
}
|
||||
}
|
||||
|
||||
fn provider_config_field_value(
|
||||
config: &Config,
|
||||
key: &crate::providers::base::ConfigKey,
|
||||
secrets: Option<&HashMap<String, serde_json::Value>>,
|
||||
) -> ProviderConfigFieldValueDto {
|
||||
let value = if key.secret {
|
||||
std::env::var(key.name.to_uppercase()).ok().or_else(|| {
|
||||
secrets
|
||||
.and_then(|values| values.get(&key.name))
|
||||
.and_then(config_value_to_string)
|
||||
})
|
||||
} else {
|
||||
config
|
||||
.get_param::<serde_json::Value>(&key.name)
|
||||
.ok()
|
||||
.and_then(|value| config_value_to_string(&value))
|
||||
};
|
||||
|
||||
ProviderConfigFieldValueDto {
|
||||
key: key.name.clone(),
|
||||
value: value.as_deref().map(|value| {
|
||||
if key.secret {
|
||||
mask_secret_value(value)
|
||||
} else {
|
||||
value.to_string()
|
||||
}
|
||||
}),
|
||||
is_set: value.is_some(),
|
||||
is_secret: key.secret,
|
||||
required: key.required,
|
||||
}
|
||||
}
|
||||
|
||||
fn refresh_skip_reason_to_dto(reason: RefreshSkipReason) -> RefreshProviderInventorySkipReasonDto {
|
||||
match reason {
|
||||
RefreshSkipReason::UnknownProvider => {
|
||||
RefreshProviderInventorySkipReasonDto::UnknownProvider
|
||||
}
|
||||
RefreshSkipReason::NotConfigured => RefreshProviderInventorySkipReasonDto::NotConfigured,
|
||||
RefreshSkipReason::DoesNotSupportRefresh => {
|
||||
RefreshProviderInventorySkipReasonDto::DoesNotSupportRefresh
|
||||
}
|
||||
RefreshSkipReason::AlreadyRefreshing => {
|
||||
RefreshProviderInventorySkipReasonDto::AlreadyRefreshing
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn refresh_plan_to_response(refresh_plan: RefreshPlan) -> RefreshProviderInventoryResponse {
|
||||
RefreshProviderInventoryResponse {
|
||||
started: refresh_plan.started,
|
||||
skipped: refresh_plan
|
||||
.skipped
|
||||
.into_iter()
|
||||
.map(|entry| RefreshProviderInventorySkipDto {
|
||||
provider_id: entry.provider_id,
|
||||
reason: refresh_skip_reason_to_dto(entry.reason),
|
||||
})
|
||||
.collect(),
|
||||
}
|
||||
}
|
||||
|
||||
impl GooseAcpAgent {
|
||||
pub(super) async fn on_list_providers(
|
||||
&self,
|
||||
req: ListProvidersRequest,
|
||||
) -> Result<ListProvidersResponse, sacp::Error> {
|
||||
let entries = self
|
||||
.provider_inventory
|
||||
.entries(&req.provider_ids)
|
||||
.await
|
||||
.internal_err()?;
|
||||
Ok(ListProvidersResponse {
|
||||
entries: entries.into_iter().map(inventory_entry_to_dto).collect(),
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) async fn provider_config_status(provider_id: String) -> ProviderConfigStatusDto {
|
||||
let is_configured = match crate::providers::get_from_registry(&provider_id).await {
|
||||
Ok(entry) => {
|
||||
match tokio::task::spawn_blocking(move || entry.inventory_configured()).await {
|
||||
Ok(is_configured) => is_configured,
|
||||
Err(error) => {
|
||||
warn!(
|
||||
provider = %provider_id,
|
||||
error = %error,
|
||||
"provider config status check failed"
|
||||
);
|
||||
false
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(_) => false,
|
||||
};
|
||||
|
||||
ProviderConfigStatusDto {
|
||||
provider_id,
|
||||
is_configured,
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) async fn provider_config_statuses(
|
||||
provider_ids: &[String],
|
||||
) -> Vec<ProviderConfigStatusDto> {
|
||||
let mut ids = if provider_ids.is_empty() {
|
||||
crate::providers::providers()
|
||||
.await
|
||||
.into_iter()
|
||||
.map(|(metadata, _)| metadata.name)
|
||||
.collect::<Vec<_>>()
|
||||
} else {
|
||||
provider_ids.to_vec()
|
||||
};
|
||||
ids.sort();
|
||||
ids.dedup();
|
||||
|
||||
let mut statuses = stream::iter(ids)
|
||||
.map(Self::provider_config_status)
|
||||
.buffer_unordered(PROVIDER_CONFIG_STATUS_CHECK_CONCURRENCY)
|
||||
.collect::<Vec<_>>()
|
||||
.await;
|
||||
statuses.sort_by(|a, b| a.provider_id.cmp(&b.provider_id));
|
||||
statuses
|
||||
}
|
||||
|
||||
pub(super) fn spawn_provider_inventory_refresh_jobs(&self, refresh_plan: &RefreshJobPlan) {
|
||||
for refresh_job in refresh_plan.started.iter().cloned() {
|
||||
let provider_inventory = self.provider_inventory.clone();
|
||||
let provider_factory = Arc::clone(&self.provider_factory);
|
||||
let provider_id = refresh_job.provider_id.clone();
|
||||
let identity = refresh_job.identity.clone();
|
||||
tokio::spawn(async move {
|
||||
let mut refresh_guard = provider_inventory.refresh_guard(&identity);
|
||||
let provider_result = AssertUnwindSafe(async {
|
||||
let metadata = crate::providers::get_from_registry(&provider_id).await?;
|
||||
let model_config =
|
||||
crate::model::ModelConfig::new(&metadata.metadata().default_model)?
|
||||
.with_canonical_limits(&provider_id);
|
||||
provider_factory(provider_id.clone(), model_config, Vec::new()).await
|
||||
})
|
||||
.catch_unwind()
|
||||
.await;
|
||||
|
||||
let fetch_result: Result<Vec<String>> = match provider_result {
|
||||
Ok(Ok(provider)) => {
|
||||
match ensure_refresh_identity_current(&provider_id, &identity).await {
|
||||
Ok(()) => match AssertUnwindSafe(provider.fetch_recommended_models())
|
||||
.catch_unwind()
|
||||
.await
|
||||
{
|
||||
Ok(Ok(models)) => Ok(models),
|
||||
Ok(Err(error)) => Err(anyhow::anyhow!(error.to_string())),
|
||||
Err(_) => {
|
||||
Err(anyhow::anyhow!("provider inventory refresh task panicked"))
|
||||
}
|
||||
},
|
||||
Err(error) => Err(error),
|
||||
}
|
||||
}
|
||||
Ok(Err(error)) => Err(error),
|
||||
Err(_) => Err(anyhow::anyhow!("provider inventory refresh task panicked")),
|
||||
};
|
||||
|
||||
match fetch_result {
|
||||
Ok(models) => match provider_inventory
|
||||
.store_refreshed_models_for_identity(&identity, &models)
|
||||
.await
|
||||
{
|
||||
Ok(()) => refresh_guard.complete(),
|
||||
Err(error) => warn!(
|
||||
provider = %provider_id,
|
||||
error = %error,
|
||||
"failed to store refreshed provider inventory"
|
||||
),
|
||||
},
|
||||
Err(error) => {
|
||||
let error_message = error.to_string();
|
||||
match provider_inventory
|
||||
.store_refresh_error_for_identity(&identity, error_message.clone())
|
||||
.await
|
||||
{
|
||||
Ok(()) => refresh_guard.complete(),
|
||||
Err(store_error) => warn!(
|
||||
provider = %provider_id,
|
||||
error = %store_error,
|
||||
refresh_error = %error_message,
|
||||
"failed to store provider inventory refresh error"
|
||||
),
|
||||
}
|
||||
warn!(provider = %provider_id, error = %error_message, "provider inventory refresh failed");
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) async fn start_provider_inventory_refresh(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<RefreshProviderInventoryResponse, sacp::Error> {
|
||||
let refresh_job_plan = self
|
||||
.provider_inventory
|
||||
.plan_refresh_jobs(provider_ids)
|
||||
.await
|
||||
.internal_err()?;
|
||||
self.spawn_provider_inventory_refresh_jobs(&refresh_job_plan);
|
||||
Ok(refresh_plan_to_response(
|
||||
refresh_job_plan.into_public_plan(),
|
||||
))
|
||||
}
|
||||
|
||||
pub(super) async fn on_refresh_provider_inventory(
|
||||
&self,
|
||||
req: RefreshProviderInventoryRequest,
|
||||
) -> Result<RefreshProviderInventoryResponse, sacp::Error> {
|
||||
Config::global().invalidate_secrets_cache();
|
||||
self.start_provider_inventory_refresh(&req.provider_ids)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(super) async fn on_read_provider_config(
|
||||
&self,
|
||||
req: ProviderConfigReadRequest,
|
||||
) -> Result<ProviderConfigReadResponse, sacp::Error> {
|
||||
let entry = crate::providers::get_from_registry(&req.provider_id)
|
||||
.await
|
||||
.invalid_params_err_ctx("Unknown provider")?;
|
||||
let config = Config::global();
|
||||
let config_keys = &entry.metadata().config_keys;
|
||||
let secrets = if config_keys.iter().any(|key| key.secret) {
|
||||
Some(config.all_secrets().internal_err()?)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
Ok(ProviderConfigReadResponse {
|
||||
fields: config_keys
|
||||
.iter()
|
||||
.map(|key| provider_config_field_value(config, key, secrets.as_ref()))
|
||||
.collect(),
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) async fn on_provider_config_status(
|
||||
&self,
|
||||
req: ProviderConfigStatusRequest,
|
||||
) -> Result<ProviderConfigStatusResponse, sacp::Error> {
|
||||
Ok(ProviderConfigStatusResponse {
|
||||
statuses: Self::provider_config_statuses(&req.provider_ids).await,
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) async fn on_save_provider_config(
|
||||
&self,
|
||||
req: ProviderConfigSaveRequest,
|
||||
) -> Result<ProviderConfigChangeResponse, sacp::Error> {
|
||||
let entry = crate::providers::get_from_registry(&req.provider_id)
|
||||
.await
|
||||
.invalid_params_err_ctx("Unknown provider")?;
|
||||
let metadata = entry.metadata().clone();
|
||||
let config = Config::global();
|
||||
let mut config_updates = Vec::new();
|
||||
let mut secret_updates = Vec::new();
|
||||
|
||||
for field in &req.fields {
|
||||
let Some(config_key) = metadata
|
||||
.config_keys
|
||||
.iter()
|
||||
.find(|config_key| config_key.name == field.key)
|
||||
else {
|
||||
return Err(sacp::Error::invalid_params()
|
||||
.data(format!("Unsupported provider config field: {}", field.key)));
|
||||
};
|
||||
|
||||
let value = field.value.trim();
|
||||
if value.is_empty() {
|
||||
return Err(sacp::Error::invalid_params().data(format!(
|
||||
"Provider config field cannot be empty: {}",
|
||||
field.key
|
||||
)));
|
||||
}
|
||||
|
||||
if config_key.secret {
|
||||
secret_updates.push((
|
||||
config_key.name.clone(),
|
||||
serde_json::Value::String(value.to_string()),
|
||||
));
|
||||
} else {
|
||||
config_updates.push((config_key.name.clone(), value.to_string()));
|
||||
}
|
||||
}
|
||||
|
||||
for (key, value) in config_updates {
|
||||
config
|
||||
.set_param(&key, &value)
|
||||
.internal_err_ctx("Failed to save provider config field")?;
|
||||
}
|
||||
config
|
||||
.set_secret_values(&secret_updates)
|
||||
.internal_err_ctx("Failed to save provider secret fields")?;
|
||||
|
||||
let provider_ids = [req.provider_id.clone()];
|
||||
let status = Self::provider_config_status(req.provider_id.clone()).await;
|
||||
let refresh = self.start_provider_inventory_refresh(&provider_ids).await?;
|
||||
Ok(ProviderConfigChangeResponse { status, refresh })
|
||||
}
|
||||
|
||||
pub(super) async fn on_delete_provider_config(
|
||||
&self,
|
||||
req: ProviderConfigDeleteRequest,
|
||||
) -> Result<ProviderConfigChangeResponse, sacp::Error> {
|
||||
let entry = crate::providers::get_from_registry(&req.provider_id)
|
||||
.await
|
||||
.invalid_params_err_ctx("Unknown provider")?;
|
||||
let metadata = entry.metadata().clone();
|
||||
let config = Config::global();
|
||||
let mut secret_keys = Vec::new();
|
||||
|
||||
for config_key in &metadata.config_keys {
|
||||
if config_key.secret {
|
||||
secret_keys.push(config_key.name.clone());
|
||||
} else {
|
||||
config
|
||||
.delete(&config_key.name)
|
||||
.internal_err_ctx("Failed to delete provider config field")?;
|
||||
}
|
||||
}
|
||||
|
||||
config
|
||||
.delete_secret_values(&secret_keys)
|
||||
.internal_err_ctx("Failed to delete provider secret fields")?;
|
||||
crate::providers::cleanup_provider(&req.provider_id)
|
||||
.await
|
||||
.internal_err_ctx("Failed to clean up provider state")?;
|
||||
|
||||
let provider_ids = [req.provider_id.clone()];
|
||||
let status = Self::provider_config_status(req.provider_id.clone()).await;
|
||||
let refresh = self.start_provider_inventory_refresh(&provider_ids).await?;
|
||||
Ok(ProviderConfigChangeResponse { status, refresh })
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
use super::*;
|
||||
|
||||
impl GooseAcpAgent {
|
||||
pub(super) async fn on_read_resource(
|
||||
&self,
|
||||
req: ReadResourceRequest,
|
||||
) -> Result<ReadResourceResponse, sacp::Error> {
|
||||
let internal_id = self.internal_session_id(&req.session_id).await?;
|
||||
let agent = self.get_session_agent(&req.session_id, None).await?;
|
||||
let cancel_token = CancellationToken::new();
|
||||
let result = agent
|
||||
.extension_manager
|
||||
.read_resource(&internal_id, &req.uri, &req.extension_name, cancel_token)
|
||||
.await
|
||||
.internal_err()?;
|
||||
let result_json = serde_json::to_value(&result).internal_err()?;
|
||||
Ok(ReadResourceResponse {
|
||||
result: result_json,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,32 @@
|
||||
use super::*;
|
||||
|
||||
impl GooseAcpAgent {
|
||||
pub(super) async fn on_check_secret(
|
||||
&self,
|
||||
req: CheckSecretRequest,
|
||||
) -> Result<CheckSecretResponse, sacp::Error> {
|
||||
let config = self.config()?;
|
||||
let exists = config.get_secret::<serde_json::Value>(&req.key).is_ok();
|
||||
Ok(CheckSecretResponse { exists })
|
||||
}
|
||||
|
||||
pub(super) async fn on_upsert_secret(
|
||||
&self,
|
||||
req: UpsertSecretRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
let config = self.config()?;
|
||||
config.set_secret(&req.key, &req.value).internal_err()?;
|
||||
Config::global().invalidate_secrets_cache();
|
||||
Ok(EmptyResponse {})
|
||||
}
|
||||
|
||||
pub(super) async fn on_remove_secret(
|
||||
&self,
|
||||
req: RemoveSecretRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
let config = self.config()?;
|
||||
config.delete_secret(&req.key).internal_err()?;
|
||||
Config::global().invalidate_secrets_cache();
|
||||
Ok(EmptyResponse {})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,175 @@
|
||||
use super::*;
|
||||
|
||||
impl GooseAcpAgent {
|
||||
pub(super) async fn on_update_working_dir(
|
||||
&self,
|
||||
req: UpdateWorkingDirRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
let working_dir = req.working_dir.trim().to_string();
|
||||
if working_dir.is_empty() {
|
||||
return Err(sacp::Error::invalid_params().data("working directory cannot be empty"));
|
||||
}
|
||||
let path = std::path::PathBuf::from(&working_dir);
|
||||
if !path.exists() || !path.is_dir() {
|
||||
return Err(sacp::Error::invalid_params().data("invalid directory path"));
|
||||
}
|
||||
let internal_id = self.internal_session_id(&req.session_id).await?;
|
||||
self.session_manager
|
||||
.update(&internal_id)
|
||||
.working_dir(path.clone())
|
||||
.apply()
|
||||
.await
|
||||
.internal_err()?;
|
||||
|
||||
self.thread_manager
|
||||
.update_working_dir(&req.session_id, &working_dir)
|
||||
.await
|
||||
.internal_err()?;
|
||||
|
||||
if let Some(session) = self.sessions.lock().await.get_mut(&req.session_id) {
|
||||
match &session.agent {
|
||||
AgentHandle::Ready(agent) => {
|
||||
agent.extension_manager.update_working_dir(&path).await;
|
||||
}
|
||||
AgentHandle::Loading(_) => {
|
||||
session.pending_working_dir = Some(path);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Ok(EmptyResponse {})
|
||||
}
|
||||
|
||||
pub(super) async fn on_delete_session(
|
||||
&self,
|
||||
req: DeleteSessionRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
// Delete the thread and all its internal sessions + messages.
|
||||
self.thread_manager
|
||||
.delete_thread(&req.session_id)
|
||||
.await
|
||||
.internal_err()?;
|
||||
self.sessions.lock().await.remove(&req.session_id);
|
||||
Ok(EmptyResponse {})
|
||||
}
|
||||
|
||||
pub(super) async fn on_export_session(
|
||||
&self,
|
||||
req: ExportSessionRequest,
|
||||
) -> Result<ExportSessionResponse, sacp::Error> {
|
||||
let thread = self
|
||||
.thread_manager
|
||||
.get_thread(&req.session_id)
|
||||
.await
|
||||
.internal_err()?;
|
||||
let internal_id = thread
|
||||
.current_session_id
|
||||
.ok_or_else(|| sacp::Error::internal_error().data("Thread has no internal session"))?;
|
||||
let data = self
|
||||
.session_manager
|
||||
.export_session(&internal_id)
|
||||
.await
|
||||
.internal_err()?;
|
||||
Ok(ExportSessionResponse { data })
|
||||
}
|
||||
|
||||
pub(super) async fn on_import_session(
|
||||
&self,
|
||||
req: ImportSessionRequest,
|
||||
) -> Result<ImportSessionResponse, sacp::Error> {
|
||||
let session = self
|
||||
.session_manager
|
||||
.import_session(&req.data, Some(SessionType::Acp))
|
||||
.await
|
||||
.internal_err()?;
|
||||
|
||||
// Create a thread for the imported session.
|
||||
let thread = self
|
||||
.thread_manager
|
||||
.create_thread(
|
||||
Some(session.name.clone()),
|
||||
None,
|
||||
Some(session.working_dir.display().to_string()),
|
||||
)
|
||||
.await
|
||||
.internal_err()?;
|
||||
|
||||
// Link the internal session to the thread.
|
||||
self.session_manager
|
||||
.update(&session.id)
|
||||
.thread_id(Some(thread.id.clone()))
|
||||
.apply()
|
||||
.await
|
||||
.internal_err()?;
|
||||
|
||||
// Copy conversation messages into thread_messages so they appear in the thread.
|
||||
if let Some(ref conversation) = session.conversation {
|
||||
for msg in conversation.messages() {
|
||||
self.thread_manager
|
||||
.append_message(&thread.id, Some(&session.id), msg)
|
||||
.await
|
||||
.internal_err()?;
|
||||
}
|
||||
}
|
||||
|
||||
// Re-fetch thread to get accurate message_count.
|
||||
let thread = self
|
||||
.thread_manager
|
||||
.get_thread(&thread.id)
|
||||
.await
|
||||
.internal_err()?;
|
||||
|
||||
Ok(ImportSessionResponse {
|
||||
session_id: thread.id,
|
||||
title: Some(thread.name),
|
||||
updated_at: Some(thread.updated_at.to_rfc3339()),
|
||||
message_count: thread.message_count as u64,
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) async fn on_update_session_project(
|
||||
&self,
|
||||
req: UpdateSessionProjectRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
let project_id = req.project_id;
|
||||
self.update_thread_metadata(&req.session_id, move |meta| {
|
||||
meta.project_id = project_id;
|
||||
})
|
||||
.await?;
|
||||
Ok(EmptyResponse {})
|
||||
}
|
||||
|
||||
pub(super) async fn on_rename_session(
|
||||
&self,
|
||||
req: RenameSessionRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
self.thread_manager
|
||||
.update_thread(&req.session_id, Some(req.title), Some(true), None)
|
||||
.await
|
||||
.map_err(|e| sacp::Error::internal_error().data(e.to_string()))?;
|
||||
Ok(EmptyResponse {})
|
||||
}
|
||||
|
||||
pub(super) async fn on_archive_session(
|
||||
&self,
|
||||
req: ArchiveSessionRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
self.thread_manager
|
||||
.archive_thread(&req.session_id)
|
||||
.await
|
||||
.internal_err()?;
|
||||
self.sessions.lock().await.remove(&req.session_id);
|
||||
Ok(EmptyResponse {})
|
||||
}
|
||||
|
||||
pub(super) async fn on_unarchive_session(
|
||||
&self,
|
||||
req: UnarchiveSessionRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
self.thread_manager
|
||||
.unarchive_thread(&req.session_id)
|
||||
.await
|
||||
.internal_err()?;
|
||||
Ok(EmptyResponse {})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,65 @@
|
||||
use super::*;
|
||||
|
||||
impl GooseAcpAgent {
|
||||
pub(super) async fn on_create_source(
|
||||
&self,
|
||||
req: CreateSourceRequest,
|
||||
) -> Result<CreateSourceResponse, sacp::Error> {
|
||||
let source = crate::sources::create_source(
|
||||
req.source_type,
|
||||
&req.name,
|
||||
&req.description,
|
||||
&req.content,
|
||||
req.global,
|
||||
req.project_dir.as_deref(),
|
||||
)?;
|
||||
Ok(CreateSourceResponse { source })
|
||||
}
|
||||
|
||||
pub(super) async fn on_list_sources(
|
||||
&self,
|
||||
req: ListSourcesRequest,
|
||||
) -> Result<ListSourcesResponse, sacp::Error> {
|
||||
let sources = crate::sources::list_sources(req.source_type, req.project_dir.as_deref())?;
|
||||
Ok(ListSourcesResponse { sources })
|
||||
}
|
||||
|
||||
pub(super) async fn on_update_source(
|
||||
&self,
|
||||
req: UpdateSourceRequest,
|
||||
) -> Result<UpdateSourceResponse, sacp::Error> {
|
||||
let source = crate::sources::update_source(
|
||||
req.source_type,
|
||||
&req.path,
|
||||
&req.name,
|
||||
&req.description,
|
||||
&req.content,
|
||||
)?;
|
||||
Ok(UpdateSourceResponse { source })
|
||||
}
|
||||
|
||||
pub(super) async fn on_delete_source(
|
||||
&self,
|
||||
req: DeleteSourceRequest,
|
||||
) -> Result<EmptyResponse, sacp::Error> {
|
||||
crate::sources::delete_source(req.source_type, &req.path)?;
|
||||
Ok(EmptyResponse {})
|
||||
}
|
||||
|
||||
pub(super) async fn on_export_source(
|
||||
&self,
|
||||
req: ExportSourceRequest,
|
||||
) -> Result<ExportSourceResponse, sacp::Error> {
|
||||
let (json, filename) = crate::sources::export_source(req.source_type, &req.path)?;
|
||||
Ok(ExportSourceResponse { json, filename })
|
||||
}
|
||||
|
||||
pub(super) async fn on_import_sources(
|
||||
&self,
|
||||
req: ImportSourcesRequest,
|
||||
) -> Result<ImportSourcesResponse, sacp::Error> {
|
||||
let sources =
|
||||
crate::sources::import_sources(&req.data, req.global, req.project_dir.as_deref())?;
|
||||
Ok(ImportSourcesResponse { sources })
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,18 @@
|
||||
use super::*;
|
||||
|
||||
impl GooseAcpAgent {
|
||||
pub(super) async fn on_get_tools(
|
||||
&self,
|
||||
req: GetToolsRequest,
|
||||
) -> Result<GetToolsResponse, sacp::Error> {
|
||||
let internal_id = self.internal_session_id(&req.session_id).await?;
|
||||
let agent = self.get_session_agent(&req.session_id, None).await?;
|
||||
let tools = agent.list_tools(&internal_id, None).await;
|
||||
let tools_json = tools
|
||||
.into_iter()
|
||||
.map(|t| serde_json::to_value(&t))
|
||||
.collect::<Result<Vec<_>, _>>()
|
||||
.internal_err()?;
|
||||
Ok(GetToolsResponse { tools: tools_json })
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user