Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
346 changes: 346 additions & 0 deletions crates/goose-sdk-types/src/custom_requests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1771,6 +1771,352 @@ pub struct ProviderInventoryEntryDto {
pub model_selection_hint: Option<String>,
}

#[derive(Debug, Default, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, JsonSchema)]
#[serde(rename_all = "snake_case")]
pub enum LocalInferenceToolCallingMode {
#[default]
Auto,
ForceNative,
ForceEmulated,
}

#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum LocalInferenceChatTemplate {
#[default]
Embedded,
Builtin {
name: String,
},
CustomInline {
template: String,
},
}

#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
#[serde(tag = "type", rename_all_fields = "camelCase")]
pub enum LocalInferenceSamplingConfig {
Greedy,
Temperature {
temperature: f32,
top_k: i32,
top_p: f32,
min_p: f32,
#[serde(default, skip_serializing_if = "Option::is_none")]
seed: Option<u32>,
},
MirostatV2 {
tau: f32,
eta: f32,
#[serde(default, skip_serializing_if = "Option::is_none")]
seed: Option<u32>,
},
}

impl Default for LocalInferenceSamplingConfig {
fn default() -> Self {
Self::Temperature {
temperature: 0.8,
top_k: 40,
top_p: 0.95,
min_p: 0.05,
seed: None,
}
}
}

#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema)]
#[serde(rename_all = "camelCase")]
pub struct LocalInferenceModelSettingsDto {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub backend_id: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub context_size: Option<u32>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub max_output_tokens: Option<usize>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub draft_model: Option<String>,
#[serde(default)]
pub sampling: LocalInferenceSamplingConfig,
pub repeat_penalty: f32,
pub repeat_last_n: i32,
pub frequency_penalty: f32,
pub presence_penalty: f32,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub n_batch: Option<u32>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub n_gpu_layers: Option<u32>,
pub use_mlock: bool,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub flash_attention: Option<bool>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub n_threads: Option<i32>,
#[serde(default)]
pub tool_calling: LocalInferenceToolCallingMode,
#[serde(default)]
pub chat_template: LocalInferenceChatTemplate,
pub enable_thinking: bool,
pub vision_capable: bool,
pub image_token_estimate: usize,
pub mmproj_size_bytes: u64,
}

#[derive(Debug, Default, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, JsonSchema)]
pub enum LocalInferenceDownloadState {
#[default]
NotDownloaded,
Downloading,
Downloaded,
}

#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema)]
#[serde(rename_all = "camelCase")]
pub struct LocalInferenceModelDownloadStatusDto {
pub state: LocalInferenceDownloadState,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub progress_percent: Option<f32>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub bytes_downloaded: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub total_bytes: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub speed_bps: Option<u64>,
}

#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema)]
#[serde(rename_all = "camelCase")]
pub struct LocalInferenceDownloadProgressDto {
pub model_id: String,
pub status: String,
pub bytes_downloaded: u64,
pub total_bytes: u64,
pub progress_percent: f32,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub speed_bps: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub eta_seconds: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub error: Option<String>,
pub task_exited: bool,
}

#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema)]
#[serde(rename_all = "camelCase")]
pub struct LocalInferenceModelDto {
pub id: String,
pub repo_id: String,
pub filename: String,
pub quantization: String,
pub size_bytes: u64,
pub status: LocalInferenceModelDownloadStatusDto,
pub recommended: bool,
pub settings: LocalInferenceModelSettingsDto,
pub vision_capable: bool,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub mmproj_status: Option<LocalInferenceModelDownloadStatusDto>,
}

#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema)]
#[serde(rename_all = "camelCase")]
pub struct LocalInferenceHfModelVariantDto {
pub variant_id: String,
pub label: String,
pub backend_id: String,
pub format: String,
pub model_id: String,
pub download_id: String,
pub size_bytes: u64,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub filename: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub download_url: Option<String>,
pub description: String,
pub quality_rank: u8,
pub sharded: bool,
pub supported: bool,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub unsupported_reason: Option<String>,
}

#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema)]
#[serde(rename_all = "camelCase")]
pub struct LocalInferenceHfGgufFileDto {
pub filename: String,
pub size_bytes: u64,
pub quantization: String,
pub download_url: String,
}

#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema)]
#[serde(rename_all = "camelCase")]
pub struct LocalInferenceHfModelInfoDto {
pub repo_id: String,
pub author: String,
pub model_name: String,
pub downloads: u64,
#[serde(default)]
pub gguf_files: Vec<LocalInferenceHfGgufFileDto>,
#[serde(default)]
pub variants: Vec<LocalInferenceHfModelVariantDto>,
}

#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)]
#[request(
method = "_goose/unstable/local-inference/models/list",
response = LocalInferenceModelsListResponse
)]
#[serde(rename_all = "camelCase")]
pub struct LocalInferenceModelsListRequest {}

#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcResponse)]
#[serde(rename_all = "camelCase")]
pub struct LocalInferenceModelsListResponse {
pub models: Vec<LocalInferenceModelDto>,
}

#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)]
#[request(
method = "_goose/unstable/local-inference/models/download",
response = LocalInferenceModelDownloadResponse
)]
#[serde(rename_all = "camelCase")]
pub struct LocalInferenceModelDownloadRequest {
pub spec: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub backend_id: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub variant_id: Option<String>,
}

#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcResponse)]
#[serde(rename_all = "camelCase")]
pub struct LocalInferenceModelDownloadResponse {
pub model_id: String,
}

#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)]
#[request(
method = "_goose/unstable/local-inference/models/download/progress",
response = LocalInferenceModelDownloadProgressResponse
)]
#[serde(rename_all = "camelCase")]
pub struct LocalInferenceModelDownloadProgressRequest {
pub model_id: String,
}

#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcResponse)]
#[serde(rename_all = "camelCase")]
pub struct LocalInferenceModelDownloadProgressResponse {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub progress: Option<LocalInferenceDownloadProgressDto>,
}

#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)]
#[request(
method = "_goose/unstable/local-inference/models/download/cancel",
response = EmptyResponse
)]
#[serde(rename_all = "camelCase")]
pub struct LocalInferenceModelDownloadCancelRequest {
pub model_id: String,
}

#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)]
#[request(
method = "_goose/unstable/local-inference/models/delete",
response = EmptyResponse
)]
#[serde(rename_all = "camelCase")]
pub struct LocalInferenceModelDeleteRequest {
pub model_id: String,
}

#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)]
#[request(
method = "_goose/unstable/local-inference/models/settings/read",
response = LocalInferenceModelSettingsReadResponse
)]
#[serde(rename_all = "camelCase")]
pub struct LocalInferenceModelSettingsReadRequest {
pub model_id: String,
}

#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcResponse)]
#[serde(rename_all = "camelCase")]
pub struct LocalInferenceModelSettingsReadResponse {
pub settings: LocalInferenceModelSettingsDto,
}

#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)]
#[request(
method = "_goose/unstable/local-inference/models/settings/update",
response = LocalInferenceModelSettingsUpdateResponse
)]
#[serde(rename_all = "camelCase")]
pub struct LocalInferenceModelSettingsUpdateRequest {
pub model_id: String,
pub settings: LocalInferenceModelSettingsDto,
}

#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcResponse)]
#[serde(rename_all = "camelCase")]
pub struct LocalInferenceModelSettingsUpdateResponse {
pub settings: LocalInferenceModelSettingsDto,
}

#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)]
#[request(
method = "_goose/unstable/local-inference/huggingface/search",
response = LocalInferenceHuggingFaceSearchResponse
)]
#[serde(rename_all = "camelCase")]
pub struct LocalInferenceHuggingFaceSearchRequest {
pub query: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub limit: Option<usize>,
}

#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcResponse)]
#[serde(rename_all = "camelCase")]
pub struct LocalInferenceHuggingFaceSearchResponse {
pub models: Vec<LocalInferenceHfModelInfoDto>,
}

#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)]
#[request(
method = "_goose/unstable/local-inference/huggingface/repo/variants",
response = LocalInferenceHuggingFaceRepoVariantsResponse
)]
#[serde(rename_all = "camelCase")]
pub struct LocalInferenceHuggingFaceRepoVariantsRequest {
pub repo_id: String,
}

#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcResponse)]
#[serde(rename_all = "camelCase")]
pub struct LocalInferenceHuggingFaceRepoVariantsResponse {
pub variants: Vec<LocalInferenceHfModelVariantDto>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub recommended_index: Option<usize>,
pub available_memory_bytes: u64,
pub downloaded_quants: Vec<String>,
pub downloaded_variants: Vec<String>,
}

#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)]
#[request(
method = "_goose/unstable/local-inference/chat-templates/builtin/list",
response = LocalInferenceBuiltinChatTemplatesListResponse
)]
#[serde(rename_all = "camelCase")]
pub struct LocalInferenceBuiltinChatTemplatesListRequest {}

#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcResponse)]
#[serde(rename_all = "camelCase")]
pub struct LocalInferenceBuiltinChatTemplatesListResponse {
pub templates: Vec<String>,
}

/// Empty success response for operations that return no data.
#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcResponse)]
pub struct EmptyResponse {}
Expand Down
27 changes: 1 addition & 26 deletions crates/goose-server/src/openapi.rs
Original file line number Diff line number Diff line change
Expand Up @@ -656,33 +656,8 @@ pub struct ApiDoc;
super::routes::dictation::get_download_progress,
super::routes::dictation::cancel_download,
super::routes::dictation::delete_model,
super::routes::local_inference::list_local_models,
super::routes::local_inference::sync_featured_models,
super::routes::local_inference::search_hf_models,
super::routes::local_inference::list_builtin_chat_templates,
super::routes::local_inference::get_repo_files,
super::routes::local_inference::download_hf_model,
super::routes::local_inference::get_local_model_download_progress,
super::routes::local_inference::cancel_local_model_download,
super::routes::local_inference::delete_local_model,
super::routes::local_inference::get_model_settings,
super::routes::local_inference::update_model_settings,
),
components(schemas(
super::routes::dictation::WhisperModelResponse,
super::routes::local_inference::LocalModelResponse,
super::routes::local_inference::ModelDownloadStatus,
super::routes::local_inference::DownloadModelRequest,
goose::providers::local_inference::hf_models::HfModelInfo,
goose::providers::local_inference::hf_models::HfModelVariant,
goose::providers::local_inference::hf_models::HfGgufFile,
goose::providers::local_inference::hf_models::HfQuantVariant,
super::routes::local_inference::RepoVariantsResponse,
goose::providers::local_inference::local_model_registry::ModelSettings,
goose::providers::local_inference::local_model_registry::ChatTemplate,
goose::providers::local_inference::local_model_registry::SamplingConfig,
goose::providers::local_inference::local_model_registry::ToolCallingMode,
))
components(schemas(super::routes::dictation::WhisperModelResponse,))
)]
pub struct LocalInferenceApiDoc;

Expand Down
Loading
Loading