added optional request params and context limit from GOOSE_PREDEFINED_MODELS (#6489)
This commit is contained in:
@@ -298,6 +298,7 @@ impl GooseAcpAgent {
|
|||||||
toolshim: false,
|
toolshim: false,
|
||||||
toolshim_model: None,
|
toolshim_model: None,
|
||||||
fast_model: None,
|
fast_model: None,
|
||||||
|
request_params: None,
|
||||||
};
|
};
|
||||||
let provider = create(&provider_name, model_config).await?;
|
let provider = create(&provider_name, model_config).await?;
|
||||||
let goose_mode = config
|
let goose_mode = config
|
||||||
|
|||||||
@@ -48,6 +48,8 @@ pub struct UpdateProviderRequest {
|
|||||||
provider: String,
|
provider: String,
|
||||||
model: Option<String>,
|
model: Option<String>,
|
||||||
session_id: String,
|
session_id: String,
|
||||||
|
context_limit: Option<usize>,
|
||||||
|
request_params: Option<std::collections::HashMap<String, serde_json::Value>>,
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Deserialize, utoipa::ToSchema)]
|
#[derive(Deserialize, utoipa::ToSchema)]
|
||||||
@@ -528,12 +530,15 @@ async fn update_agent_provider(
|
|||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
let model_config = ModelConfig::new(&model).map_err(|e| {
|
let model_config = ModelConfig::new(&model)
|
||||||
(
|
.map_err(|e| {
|
||||||
StatusCode::BAD_REQUEST,
|
(
|
||||||
format!("Invalid model config: {}", e),
|
StatusCode::BAD_REQUEST,
|
||||||
)
|
format!("Invalid model config: {}", e),
|
||||||
})?;
|
)
|
||||||
|
})?
|
||||||
|
.with_context_limit(payload.context_limit)
|
||||||
|
.with_request_params(payload.request_params);
|
||||||
|
|
||||||
let new_provider = create(&payload.provider, model_config).await.map_err(|e| {
|
let new_provider = create(&payload.provider, model_config).await.map_err(|e| {
|
||||||
(
|
(
|
||||||
|
|||||||
@@ -444,6 +444,7 @@ mod tests {
|
|||||||
toolshim: false,
|
toolshim: false,
|
||||||
toolshim_model: None,
|
toolshim_model: None,
|
||||||
fast_model: None,
|
fast_model: None,
|
||||||
|
request_params: None,
|
||||||
},
|
},
|
||||||
max_tool_responses: None,
|
max_tool_responses: None,
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,10 +1,39 @@
|
|||||||
use once_cell::sync::Lazy;
|
use once_cell::sync::Lazy;
|
||||||
use serde::{Deserialize, Serialize};
|
use serde::{Deserialize, Serialize};
|
||||||
|
use serde_json::Value;
|
||||||
|
use std::collections::HashMap;
|
||||||
use thiserror::Error;
|
use thiserror::Error;
|
||||||
use utoipa::ToSchema;
|
use utoipa::ToSchema;
|
||||||
|
|
||||||
const DEFAULT_CONTEXT_LIMIT: usize = 128_000;
|
const DEFAULT_CONTEXT_LIMIT: usize = 128_000;
|
||||||
|
|
||||||
|
#[derive(Debug, Clone, Deserialize)]
|
||||||
|
struct PredefinedModel {
|
||||||
|
name: String,
|
||||||
|
#[serde(default)]
|
||||||
|
context_limit: Option<usize>,
|
||||||
|
#[serde(default)]
|
||||||
|
request_params: Option<HashMap<String, Value>>,
|
||||||
|
}
|
||||||
|
|
||||||
|
fn get_predefined_models() -> Vec<PredefinedModel> {
|
||||||
|
static PREDEFINED_MODELS: Lazy<Vec<PredefinedModel>> =
|
||||||
|
Lazy::new(|| match std::env::var("GOOSE_PREDEFINED_MODELS") {
|
||||||
|
Ok(json_str) => serde_json::from_str(&json_str).unwrap_or_else(|e| {
|
||||||
|
tracing::warn!("Failed to parse GOOSE_PREDEFINED_MODELS: {}", e);
|
||||||
|
Vec::new()
|
||||||
|
}),
|
||||||
|
Err(_) => Vec::new(),
|
||||||
|
});
|
||||||
|
PREDEFINED_MODELS.clone()
|
||||||
|
}
|
||||||
|
|
||||||
|
fn find_predefined_model(model_name: &str) -> Option<PredefinedModel> {
|
||||||
|
get_predefined_models()
|
||||||
|
.into_iter()
|
||||||
|
.find(|m| m.name == model_name)
|
||||||
|
}
|
||||||
|
|
||||||
#[derive(Error, Debug)]
|
#[derive(Error, Debug)]
|
||||||
pub enum ConfigError {
|
pub enum ConfigError {
|
||||||
#[error("Environment variable '{0}' not found")]
|
#[error("Environment variable '{0}' not found")]
|
||||||
@@ -80,6 +109,9 @@ pub struct ModelConfig {
|
|||||||
pub toolshim: bool,
|
pub toolshim: bool,
|
||||||
pub toolshim_model: Option<String>,
|
pub toolshim_model: Option<String>,
|
||||||
pub fast_model: Option<String>,
|
pub fast_model: Option<String>,
|
||||||
|
/// Provider-specific request parameters (e.g., anthropic_beta headers)
|
||||||
|
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||||
|
pub request_params: Option<HashMap<String, Value>>,
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||||
@@ -97,7 +129,26 @@ impl ModelConfig {
|
|||||||
model_name: String,
|
model_name: String,
|
||||||
context_env_var: Option<&str>,
|
context_env_var: Option<&str>,
|
||||||
) -> Result<Self, ConfigError> {
|
) -> Result<Self, ConfigError> {
|
||||||
let context_limit = Self::parse_context_limit(&model_name, None, context_env_var)?;
|
let predefined = find_predefined_model(&model_name);
|
||||||
|
|
||||||
|
let context_limit = if let Some(ref pm) = predefined {
|
||||||
|
if let Some(env_var) = context_env_var {
|
||||||
|
if let Ok(val) = std::env::var(env_var) {
|
||||||
|
Some(Self::validate_context_limit(&val, env_var)?)
|
||||||
|
} else {
|
||||||
|
pm.context_limit
|
||||||
|
}
|
||||||
|
} else if let Ok(val) = std::env::var("GOOSE_CONTEXT_LIMIT") {
|
||||||
|
Some(Self::validate_context_limit(&val, "GOOSE_CONTEXT_LIMIT")?)
|
||||||
|
} else {
|
||||||
|
pm.context_limit
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
Self::parse_context_limit(&model_name, None, context_env_var)?
|
||||||
|
};
|
||||||
|
|
||||||
|
let request_params = predefined.and_then(|pm| pm.request_params);
|
||||||
|
|
||||||
let temperature = Self::parse_temperature()?;
|
let temperature = Self::parse_temperature()?;
|
||||||
let max_tokens = Self::parse_max_tokens()?;
|
let max_tokens = Self::parse_max_tokens()?;
|
||||||
let toolshim = Self::parse_toolshim()?;
|
let toolshim = Self::parse_toolshim()?;
|
||||||
@@ -111,6 +162,7 @@ impl ModelConfig {
|
|||||||
toolshim,
|
toolshim,
|
||||||
toolshim_model,
|
toolshim_model,
|
||||||
fast_model: None,
|
fast_model: None,
|
||||||
|
request_params,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -285,6 +337,11 @@ impl ModelConfig {
|
|||||||
self
|
self
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub fn with_request_params(mut self, params: Option<HashMap<String, Value>>) -> Self {
|
||||||
|
self.request_params = params;
|
||||||
|
self
|
||||||
|
}
|
||||||
|
|
||||||
pub fn use_fast_model(&self) -> Self {
|
pub fn use_fast_model(&self) -> Self {
|
||||||
if let Some(fast_model) = &self.fast_model {
|
if let Some(fast_model) = &self.fast_model {
|
||||||
let mut config = self.clone();
|
let mut config = self.clone();
|
||||||
|
|||||||
@@ -642,6 +642,15 @@ pub fn create_request(
|
|||||||
apply_cache_control_for_claude(&mut payload);
|
apply_cache_control_for_claude(&mut payload);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Add request_params to the payload (e.g., anthropic_beta for extended context)
|
||||||
|
if let Some(params) = &model_config.request_params {
|
||||||
|
if let Some(obj) = payload.as_object_mut() {
|
||||||
|
for (key, value) in params {
|
||||||
|
obj.insert(key.clone(), value.clone());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
Ok(payload)
|
Ok(payload)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1013,6 +1022,7 @@ mod tests {
|
|||||||
toolshim: false,
|
toolshim: false,
|
||||||
toolshim_model: None,
|
toolshim_model: None,
|
||||||
fast_model: None,
|
fast_model: None,
|
||||||
|
request_params: None,
|
||||||
};
|
};
|
||||||
let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?;
|
let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?;
|
||||||
let obj = request.as_object().unwrap();
|
let obj = request.as_object().unwrap();
|
||||||
@@ -1044,6 +1054,7 @@ mod tests {
|
|||||||
toolshim: false,
|
toolshim: false,
|
||||||
toolshim_model: None,
|
toolshim_model: None,
|
||||||
fast_model: None,
|
fast_model: None,
|
||||||
|
request_params: None,
|
||||||
};
|
};
|
||||||
let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?;
|
let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?;
|
||||||
assert_eq!(request["reasoning_effort"], "high");
|
assert_eq!(request["reasoning_effort"], "high");
|
||||||
@@ -1360,6 +1371,7 @@ mod tests {
|
|||||||
toolshim: false,
|
toolshim: false,
|
||||||
toolshim_model: None,
|
toolshim_model: None,
|
||||||
fast_model: None,
|
fast_model: None,
|
||||||
|
request_params: None,
|
||||||
};
|
};
|
||||||
|
|
||||||
let messages = vec![
|
let messages = vec![
|
||||||
@@ -1411,6 +1423,7 @@ mod tests {
|
|||||||
toolshim: false,
|
toolshim: false,
|
||||||
toolshim_model: None,
|
toolshim_model: None,
|
||||||
fast_model: None,
|
fast_model: None,
|
||||||
|
request_params: None,
|
||||||
};
|
};
|
||||||
|
|
||||||
let messages = vec![Message::user().with_text("Hello")];
|
let messages = vec![Message::user().with_text("Hello")];
|
||||||
|
|||||||
@@ -1295,6 +1295,7 @@ mod tests {
|
|||||||
toolshim: false,
|
toolshim: false,
|
||||||
toolshim_model: None,
|
toolshim_model: None,
|
||||||
fast_model: None,
|
fast_model: None,
|
||||||
|
request_params: None,
|
||||||
};
|
};
|
||||||
let request = create_request(
|
let request = create_request(
|
||||||
&model_config,
|
&model_config,
|
||||||
@@ -1334,6 +1335,7 @@ mod tests {
|
|||||||
toolshim: false,
|
toolshim: false,
|
||||||
toolshim_model: None,
|
toolshim_model: None,
|
||||||
fast_model: None,
|
fast_model: None,
|
||||||
|
request_params: None,
|
||||||
};
|
};
|
||||||
let request = create_request(
|
let request = create_request(
|
||||||
&model_config,
|
&model_config,
|
||||||
@@ -1374,6 +1376,7 @@ mod tests {
|
|||||||
toolshim: false,
|
toolshim: false,
|
||||||
toolshim_model: None,
|
toolshim_model: None,
|
||||||
fast_model: None,
|
fast_model: None,
|
||||||
|
request_params: None,
|
||||||
};
|
};
|
||||||
let request = create_request(
|
let request = create_request(
|
||||||
&model_config,
|
&model_config,
|
||||||
|
|||||||
@@ -4524,6 +4524,12 @@
|
|||||||
"model_name": {
|
"model_name": {
|
||||||
"type": "string"
|
"type": "string"
|
||||||
},
|
},
|
||||||
|
"request_params": {
|
||||||
|
"type": "object",
|
||||||
|
"description": "Provider-specific request parameters (e.g., anthropic_beta headers)",
|
||||||
|
"additionalProperties": {},
|
||||||
|
"nullable": true
|
||||||
|
},
|
||||||
"temperature": {
|
"temperature": {
|
||||||
"type": "number",
|
"type": "number",
|
||||||
"format": "float",
|
"format": "float",
|
||||||
@@ -6327,6 +6333,11 @@
|
|||||||
"session_id"
|
"session_id"
|
||||||
],
|
],
|
||||||
"properties": {
|
"properties": {
|
||||||
|
"context_limit": {
|
||||||
|
"type": "integer",
|
||||||
|
"nullable": true,
|
||||||
|
"minimum": 0
|
||||||
|
},
|
||||||
"model": {
|
"model": {
|
||||||
"type": "string",
|
"type": "string",
|
||||||
"nullable": true
|
"nullable": true
|
||||||
@@ -6334,6 +6345,11 @@
|
|||||||
"provider": {
|
"provider": {
|
||||||
"type": "string"
|
"type": "string"
|
||||||
},
|
},
|
||||||
|
"request_params": {
|
||||||
|
"type": "object",
|
||||||
|
"additionalProperties": {},
|
||||||
|
"nullable": true
|
||||||
|
},
|
||||||
"session_id": {
|
"session_id": {
|
||||||
"type": "string"
|
"type": "string"
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -542,6 +542,12 @@ export type ModelConfig = {
|
|||||||
fast_model?: string | null;
|
fast_model?: string | null;
|
||||||
max_tokens?: number | null;
|
max_tokens?: number | null;
|
||||||
model_name: string;
|
model_name: string;
|
||||||
|
/**
|
||||||
|
* Provider-specific request parameters (e.g., anthropic_beta headers)
|
||||||
|
*/
|
||||||
|
request_params?: {
|
||||||
|
[key: string]: unknown;
|
||||||
|
} | null;
|
||||||
temperature?: number | null;
|
temperature?: number | null;
|
||||||
toolshim: boolean;
|
toolshim: boolean;
|
||||||
toolshim_model?: string | null;
|
toolshim_model?: string | null;
|
||||||
@@ -1168,8 +1174,12 @@ export type UpdateFromSessionRequest = {
|
|||||||
};
|
};
|
||||||
|
|
||||||
export type UpdateProviderRequest = {
|
export type UpdateProviderRequest = {
|
||||||
|
context_limit?: number | null;
|
||||||
model?: string | null;
|
model?: string | null;
|
||||||
provider: string;
|
provider: string;
|
||||||
|
request_params?: {
|
||||||
|
[key: string]: unknown;
|
||||||
|
} | null;
|
||||||
session_id: string;
|
session_id: string;
|
||||||
};
|
};
|
||||||
|
|
||||||
|
|||||||
@@ -31,6 +31,7 @@ import { getSession, Message } from '../api';
|
|||||||
import CreateRecipeFromSessionModal from './recipes/CreateRecipeFromSessionModal';
|
import CreateRecipeFromSessionModal from './recipes/CreateRecipeFromSessionModal';
|
||||||
import CreateEditRecipeModal from './recipes/CreateEditRecipeModal';
|
import CreateEditRecipeModal from './recipes/CreateEditRecipeModal';
|
||||||
import { getInitialWorkingDir } from '../utils/workingDir';
|
import { getInitialWorkingDir } from '../utils/workingDir';
|
||||||
|
import { getPredefinedModelsFromEnv } from './settings/models/predefinedModelsUtils';
|
||||||
import {
|
import {
|
||||||
trackFileAttached,
|
trackFileAttached,
|
||||||
trackVoiceDictation,
|
trackVoiceDictation,
|
||||||
@@ -448,6 +449,15 @@ export default function ChatInput({
|
|||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// First, check predefined models from environment (highest priority)
|
||||||
|
const predefinedModels = getPredefinedModelsFromEnv();
|
||||||
|
const predefinedModel = predefinedModels.find((m) => m.name === model);
|
||||||
|
if (predefinedModel?.context_limit) {
|
||||||
|
setTokenLimit(predefinedModel.context_limit);
|
||||||
|
setIsTokenLimitLoaded(true);
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
const providers = await getProviders(true);
|
const providers = await getProviders(true);
|
||||||
|
|
||||||
// Find the provider details for the current provider
|
// Find the provider details for the current provider
|
||||||
|
|||||||
@@ -53,6 +53,8 @@ export const ModelAndProviderProvider: React.FC<ModelAndProviderProviderProps> =
|
|||||||
session_id: sessionId,
|
session_id: sessionId,
|
||||||
provider: providerName,
|
provider: providerName,
|
||||||
model: modelName,
|
model: modelName,
|
||||||
|
context_limit: model.context_limit,
|
||||||
|
request_params: model.request_params,
|
||||||
},
|
},
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -24,6 +24,17 @@ const alertStyles: Record<AlertType, string> = {
|
|||||||
[AlertType.Info]: 'dark:bg-white dark:text-black bg-black text-white',
|
[AlertType.Info]: 'dark:bg-white dark:text-black bg-black text-white',
|
||||||
};
|
};
|
||||||
|
|
||||||
|
const formatTokenCount = (count: number): string => {
|
||||||
|
if (count >= 1000000) {
|
||||||
|
const millions = count / 1000000;
|
||||||
|
return millions % 1 === 0 ? `${millions.toFixed(0)}M` : `${millions.toFixed(1)}M`;
|
||||||
|
} else if (count >= 1000) {
|
||||||
|
const thousands = count / 1000;
|
||||||
|
return thousands % 1 === 0 ? `${thousands.toFixed(0)}k` : `${thousands.toFixed(1)}k`;
|
||||||
|
}
|
||||||
|
return count.toString();
|
||||||
|
};
|
||||||
|
|
||||||
export const AlertBox = ({ alert, className }: AlertBoxProps) => {
|
export const AlertBox = ({ alert, className }: AlertBoxProps) => {
|
||||||
const { read } = useConfig();
|
const { read } = useConfig();
|
||||||
const [isEditingThreshold, setIsEditingThreshold] = useState(false);
|
const [isEditingThreshold, setIsEditingThreshold] = useState(false);
|
||||||
@@ -242,18 +253,14 @@ export const AlertBox = ({ alert, className }: AlertBoxProps) => {
|
|||||||
<div className="flex justify-between items-baseline text-[11px]">
|
<div className="flex justify-between items-baseline text-[11px]">
|
||||||
<div className="flex gap-1 items-baseline">
|
<div className="flex gap-1 items-baseline">
|
||||||
<span className={'dark:text-black/60 text-white/60'}>
|
<span className={'dark:text-black/60 text-white/60'}>
|
||||||
{alert.progress!.current >= 1000
|
{formatTokenCount(alert.progress!.current)}
|
||||||
? (alert.progress!.current / 1000).toFixed(1) + 'k'
|
|
||||||
: alert.progress!.current}
|
|
||||||
</span>
|
</span>
|
||||||
<span className={'dark:text-black/40 text-white/40'}>
|
<span className={'dark:text-black/40 text-white/40'}>
|
||||||
{Math.round((alert.progress!.current / alert.progress!.total) * 100)}%
|
{Math.round((alert.progress!.current / alert.progress!.total) * 100)}%
|
||||||
</span>
|
</span>
|
||||||
</div>
|
</div>
|
||||||
<span className={'dark:text-black/60 text-white/60'}>
|
<span className={'dark:text-black/60 text-white/60'}>
|
||||||
{alert.progress!.total >= 1000
|
{formatTokenCount(alert.progress!.total)}
|
||||||
? (alert.progress!.total / 1000).toFixed(0) + 'k'
|
|
||||||
: alert.progress!.total}
|
|
||||||
</span>
|
</span>
|
||||||
</div>
|
</div>
|
||||||
{alert.showCompactButton && alert.onCompact && (
|
{alert.showCompactButton && alert.onCompact && (
|
||||||
|
|||||||
@@ -7,6 +7,8 @@ export default interface Model {
|
|||||||
lastUsed?: string;
|
lastUsed?: string;
|
||||||
alias?: string; // optional model display name
|
alias?: string; // optional model display name
|
||||||
subtext?: string; // goes below model name if not the provider
|
subtext?: string; // goes below model name if not the provider
|
||||||
|
context_limit?: number; // optional context limit override
|
||||||
|
request_params?: Record<string, unknown>; // provider-specific request parameters
|
||||||
}
|
}
|
||||||
|
|
||||||
export function createModelStruct(
|
export function createModelStruct(
|
||||||
|
|||||||
Reference in New Issue
Block a user