feat: Add lead-worker model selection and real-time model display in GUI (#2964)

Co-authored-by: jack <jack@deck.local>
This commit is contained in:
jack
2025-06-18 05:40:20 +02:00
committed by GitHub
parent d8b6e6011b
commit 657718d8c0
16 changed files with 475 additions and 9 deletions
+4
View File
@@ -589,6 +589,10 @@ async fn process_message_streaming(
// For now, we'll just log them
tracing::info!("Received MCP notification in web interface");
}
Ok(AgentEvent::ModelChange { model, mode }) => {
// Log model change
tracing::info!("Model changed to {} in {} mode", model, mode);
}
Err(e) => {
error!("Error in message stream: {}", e);
let mut sender = sender.lock().await;
+6
View File
@@ -928,6 +928,12 @@ impl Session {
}
}
}
Some(Ok(AgentEvent::ModelChange { model, mode })) => {
// Log model change if in debug mode
if self.debug {
eprintln!("Model changed to {} in {} mode", model, mode);
}
}
Some(Err(e)) => {
eprintln!("Error: {}", e);
drop(stream);
+3
View File
@@ -266,6 +266,9 @@ pub unsafe extern "C" fn goose_agent_send_message(
Ok(AgentEvent::McpNotification(_)) => {
// TODO: Handle MCP notifications.
}
Ok(AgentEvent::ModelChange { .. }) => {
// Model change events are informational, just continue
}
Err(e) => {
full_response.push_str(&format!("\nError in message stream: {}", e));
}
@@ -165,6 +165,9 @@ async fn execute_recipe(job_id: &str, recipe_path: &str) -> Result<String> {
Ok(AgentEvent::McpNotification(_)) => {
// Handle notifications if needed
}
Ok(AgentEvent::ModelChange { .. }) => {
// Model change events are informational, just continue
}
Err(e) => {
return Err(anyhow!("Error receiving message from agent: {}", e));
}
@@ -441,6 +441,26 @@ pub async fn backup_config(
}
}
#[utoipa::path(
get,
path = "/config/current-model",
responses(
(status = 200, description = "Current model retrieved successfully", body = String),
)
)]
pub async fn get_current_model(
State(state): State<Arc<AppState>>,
headers: HeaderMap,
) -> Result<Json<Value>, StatusCode> {
verify_secret_key(&headers, &state)?;
let current_model = goose::providers::base::get_current_model();
Ok(Json(serde_json::json!({
"model": current_model
})))
}
pub fn routes(state: Arc<AppState>) -> Router {
Router::new()
.route("/config", get(read_all_config))
@@ -454,6 +474,7 @@ pub fn routes(state: Arc<AppState>) -> Router {
.route("/config/init", post(init_config))
.route("/config/backup", post(backup_config))
.route("/config/permissions", post(upsert_permissions))
.route("/config/current-model", get(get_current_model))
.with_state(state)
}
+19
View File
@@ -88,6 +88,10 @@ enum MessageEvent {
Finish {
reason: String,
},
ModelChange {
model: String,
mode: String,
},
Notification {
request_id: String,
message: JsonRpcMessage,
@@ -233,6 +237,17 @@ async fn handler(
}
});
}
Ok(Some(Ok(AgentEvent::ModelChange { model, mode }))) => {
if let Err(e) = stream_event(MessageEvent::ModelChange { model, mode }, &tx).await {
tracing::error!("Error sending model change through channel: {}", e);
let _ = stream_event(
MessageEvent::Error {
error: e.to_string(),
},
&tx,
).await;
}
}
Ok(Some(Ok(AgentEvent::McpNotification((request_id, n))))) => {
if let Err(e) = stream_event(MessageEvent::Notification{
request_id: request_id.clone(),
@@ -352,6 +367,10 @@ async fn ask_handler(
}
}
}
Ok(AgentEvent::ModelChange { model, mode }) => {
// Log model change for non-streaming
tracing::info!("Model changed to {} in {} mode", model, mode);
}
Ok(AgentEvent::McpNotification(n)) => {
// Handle notifications if needed
tracing::info!("Received notification: {:?}", n);
+21
View File
@@ -65,6 +65,7 @@ pub struct Agent {
pub enum AgentEvent {
Message(Message),
McpNotification((String, JsonRpcMessage)),
ModelChange { model: String, mode: String },
}
impl Agent {
@@ -582,6 +583,26 @@ impl Agent {
&toolshim_tools,
).await {
Ok((response, usage)) => {
// Emit model change event if provider is lead-worker
let provider = self.provider().await?;
if let Some(lead_worker) = provider.as_lead_worker() {
// The actual model used is in the usage
let active_model = usage.model.clone();
let (lead_model, worker_model) = lead_worker.get_model_info();
let mode = if active_model == lead_model {
"lead"
} else if active_model == worker_model {
"worker"
} else {
"unknown"
};
yield AgentEvent::ModelChange {
model: active_model,
mode: mode.to_string(),
};
}
// record usage for the session in the session file
if let Some(session_config) = session.clone() {
Self::update_session_metrics(session_config, &usage, messages.len()).await?;
+14
View File
@@ -152,6 +152,9 @@ use async_trait::async_trait;
pub trait LeadWorkerProviderTrait {
/// Get information about the lead and worker models for logging
fn get_model_info(&self) -> (String, String);
/// Get the currently active model name
fn get_active_model(&self) -> String;
}
/// Base trait for AI providers (OpenAI, Anthropic, etc)
@@ -207,6 +210,17 @@ pub trait Provider: Send + Sync {
fn as_lead_worker(&self) -> Option<&dyn LeadWorkerProviderTrait> {
None
}
/// Get the currently active model name
/// For regular providers, this returns the configured model
/// For LeadWorkerProvider, this returns the currently active model (lead or worker)
fn get_active_model_name(&self) -> String {
if let Some(lead_worker) = self.as_lead_worker() {
lead_worker.get_active_model()
} else {
self.get_model_config().model_name
}
}
}
#[cfg(test)]
+26 -4
View File
@@ -291,6 +291,16 @@ impl LeadWorkerProviderTrait for LeadWorkerProvider {
let worker_model = self.worker_provider.get_model_config().model_name;
(lead_model, worker_model)
}
/// Get the currently active model name
fn get_active_model(&self) -> String {
// Read from the global store which was set during complete()
use super::base::get_current_model;
get_current_model().unwrap_or_else(|| {
// Fallback to lead model if no current model is set
self.lead_provider.get_model_config().model_name
})
}
}
#[async_trait]
@@ -336,19 +346,31 @@ impl Provider for LeadWorkerProvider {
"worker"
};
// Get the active model name and update the global store
let active_model_name = if turn_count < self.lead_turns || in_fallback {
self.lead_provider.get_model_config().model_name.clone()
} else {
self.worker_provider.get_model_config().model_name.clone()
};
// Update the global current model store
super::base::set_current_model(&active_model_name);
if in_fallback {
tracing::info!(
"🔄 Using {} provider for turn {} (FALLBACK MODE: {} turns remaining)",
"🔄 Using {} provider for turn {} (FALLBACK MODE: {} turns remaining) - Model: {}",
provider_type,
turn_count + 1,
fallback_remaining
fallback_remaining,
active_model_name
);
} else {
tracing::info!(
"Using {} provider for turn {} (lead_turns: {})",
"Using {} provider for turn {} (lead_turns: {}) - Model: {}",
provider_type,
turn_count + 1,
self.lead_turns
self.lead_turns,
active_model_name
);
}
+3
View File
@@ -1114,6 +1114,9 @@ async fn run_scheduled_job_internal(
Ok(AgentEvent::McpNotification(_)) => {
// Handle notifications if needed
}
Ok(AgentEvent::ModelChange { .. }) => {
// Model change events are informational, just continue
}
Err(e) => {
tracing::error!(
"[Job {}] Error receiving message from agent: {}",
+3
View File
@@ -136,6 +136,9 @@ async fn run_truncate_test(
Ok(AgentEvent::McpNotification(n)) => {
println!("MCP Notification: {n:?}");
}
Ok(AgentEvent::ModelChange { .. }) => {
// Model change events are informational, just continue
}
Err(e) => {
println!("Error: {:?}", e);
return Err(e);
+9 -2
View File
@@ -1,4 +1,4 @@
import React, { useEffect, useRef, useState, useMemo, useCallback } from 'react';
import React, { useEffect, useRef, useState, useMemo, useCallback, createContext, useContext } from 'react';
import { getApiUrl } from '../config';
import FlappyGoose from './FlappyGoose';
import GooseMessage from './GooseMessage';
@@ -37,6 +37,10 @@ import {
TextContent,
} from '../types/message';
// Context for sharing current model info
const CurrentModelContext = createContext<{ model: string; mode: string } | null>(null);
export const useCurrentModelInfo = () => useContext(CurrentModelContext);
export interface ChatType {
id: string;
title: string;
@@ -144,6 +148,7 @@ function ChatContent({
handleSubmit: _submitMessage,
updateMessageStreamBody,
notifications,
currentModelInfo,
} = useMessageStream({
api: getApiUrl('/reply'),
initialMessages: chat.messages,
@@ -504,7 +509,8 @@ function ChatContent({
}, new Map());
return (
<div className="flex flex-col w-full h-screen items-center justify-center">
<CurrentModelContext.Provider value={currentModelInfo}>
<div className="flex flex-col w-full h-screen items-center justify-center">
{/* Loader when generating recipe */}
{isGeneratingRecipe && <LayingEggLoader />}
<MoreMenuLayout
@@ -647,5 +653,6 @@ function ChatContent({
summaryContent={summaryContent}
/>
</div>
</CurrentModelContext.Provider>
);
}
@@ -2,8 +2,12 @@ import { Sliders } from 'lucide-react';
import React, { useEffect, useState, useRef } from 'react';
import { useModelAndProvider } from '../../../ModelAndProviderContext';
import { AddModelModal } from '../subcomponents/AddModelModal';
import { LeadWorkerSettings } from '../subcomponents/LeadWorkerSettings';
import { View } from '../../../../App';
import { Tooltip, TooltipTrigger, TooltipContent, TooltipProvider } from '../../../ui/Tooltip';
import Modal from '../../../Modal';
import { useCurrentModelInfo } from '../../../ChatView';
import { useConfig } from '../../../ConfigContext';
interface ModelsBottomBarProps {
dropdownRef: React.RefObject<HTMLDivElement>;
@@ -12,15 +16,36 @@ interface ModelsBottomBarProps {
export default function ModelsBottomBar({ dropdownRef, setView }: ModelsBottomBarProps) {
const { currentModel, currentProvider, getCurrentModelAndProviderForDisplay } =
useModelAndProvider();
const currentModelInfo = useCurrentModelInfo();
const { read } = useConfig();
const [isModelMenuOpen, setIsModelMenuOpen] = useState(false);
const [displayProvider, setDisplayProvider] = useState<string | null>(null);
const [isAddModelModalOpen, setIsAddModelModalOpen] = useState(false);
const [isLeadWorkerModalOpen, setIsLeadWorkerModalOpen] = useState(false);
const [isLeadWorkerActive, setIsLeadWorkerActive] = useState(false);
const menuRef = useRef<HTMLDivElement>(null);
const [isModelTruncated, setIsModelTruncated] = useState(false);
// eslint-disable-next-line no-undef
const modelRef = useRef<HTMLSpanElement>(null);
const [isTooltipOpen, setIsTooltipOpen] = useState(false);
// Check if lead/worker mode is active
useEffect(() => {
const checkLeadWorker = async () => {
try {
const leadModel = await read('GOOSE_LEAD_MODEL', false);
setIsLeadWorkerActive(!!leadModel);
} catch (error) {
setIsLeadWorkerActive(false);
}
};
checkLeadWorker();
}, [read]);
// Determine which model to display - activeModel takes priority when lead/worker is active
const displayModel = (isLeadWorkerActive && currentModelInfo?.model) ? currentModelInfo.model : (currentModel || 'Select Model');
const modelMode = currentModelInfo?.mode;
// Update display provider when current provider changes
useEffect(() => {
if (currentProvider) {
@@ -40,7 +65,7 @@ export default function ModelsBottomBar({ dropdownRef, setView }: ModelsBottomBa
checkTruncation();
window.addEventListener('resize', checkTruncation);
return () => window.removeEventListener('resize', checkTruncation);
}, [currentModel]);
}, [displayModel]);
useEffect(() => {
setIsTooltipOpen(false);
@@ -79,12 +104,22 @@ export default function ModelsBottomBar({ dropdownRef, setView }: ModelsBottomBa
ref={modelRef}
className="truncate max-w-[130px] md:max-w-[200px] lg:max-w-[360px] min-w-0 block"
>
{currentModel || 'Select Model'}
{displayModel}
{isLeadWorkerActive && modelMode && (
<span className="ml-1 text-[10px] opacity-60">
({modelMode})
</span>
)}
</span>
</TooltipTrigger>
{isModelTruncated && (
<TooltipContent className="max-w-96 overflow-auto scrollbar-thin" side="top">
{currentModel || 'Select Model'}
{displayModel}
{isLeadWorkerActive && modelMode && (
<span className="ml-1 text-[10px] opacity-60">
({modelMode})
</span>
)}
</TooltipContent>
)}
</Tooltip>
@@ -110,6 +145,17 @@ export default function ModelsBottomBar({ dropdownRef, setView }: ModelsBottomBa
<span className="text-sm">Change Model</span>
<Sliders className="w-4 h-4 ml-2 rotate-90" />
</div>
<div
className="flex items-center justify-between text-textStandard p-2 cursor-pointer transition-colors hover:bg-bgStandard
border-t border-borderSubtle"
onClick={() => {
setIsModelMenuOpen(false);
setIsLeadWorkerModalOpen(true);
}}
>
<span className="text-sm">Lead/Worker Settings</span>
<Sliders className="w-4 h-4 ml-2" />
</div>
</div>
</div>
)}
@@ -118,6 +164,12 @@ export default function ModelsBottomBar({ dropdownRef, setView }: ModelsBottomBa
{isAddModelModalOpen ? (
<AddModelModal setView={setView} onClose={() => setIsAddModelModalOpen(false)} />
) : null}
{isLeadWorkerModalOpen ? (
<Modal onClose={() => setIsLeadWorkerModalOpen(false)}>
<LeadWorkerSettings onClose={() => setIsLeadWorkerModalOpen(false)} />
</Modal>
) : null}
</div>
);
}
@@ -0,0 +1,262 @@
import { useState, useEffect } from 'react';
import { useConfig } from '../../../ConfigContext';
import { useModelAndProvider } from '../../../ModelAndProviderContext';
import { Button } from '../../../ui/button';
import { Select } from '../../../ui/Select';
import { Input } from '../../../ui/input';
import { Info } from 'lucide-react';
interface LeadWorkerSettingsProps {
onClose: () => void;
}
export function LeadWorkerSettings({ onClose }: LeadWorkerSettingsProps) {
const { read, upsert, getProviders, remove } = useConfig();
const { currentModel } = useModelAndProvider();
const [leadModel, setLeadModel] = useState<string>('');
const [workerModel, setWorkerModel] = useState<string>('');
const [leadProvider, setLeadProvider] = useState<string>('');
const [workerProvider, setWorkerProvider] = useState<string>('');
const [leadTurns, setLeadTurns] = useState<number>(3);
const [failureThreshold, setFailureThreshold] = useState<number>(2);
const [fallbackTurns, setFallbackTurns] = useState<number>(2);
const [isEnabled, setIsEnabled] = useState(false);
const [modelOptions, setModelOptions] = useState<{ value: string; label: string; provider: string }[]>([]);
const [isLoading, setIsLoading] = useState(true);
// Load current configuration
useEffect(() => {
const loadConfig = async () => {
try {
setIsLoading(true);
const [
leadModelConfig,
leadProviderConfig,
leadTurnsConfig,
failureThresholdConfig,
fallbackTurnsConfig,
] = await Promise.all([
read('GOOSE_LEAD_MODEL', false),
read('GOOSE_LEAD_PROVIDER', false),
read('GOOSE_LEAD_TURNS', false),
read('GOOSE_LEAD_FAILURE_THRESHOLD', false),
read('GOOSE_LEAD_FALLBACK_TURNS', false),
]);
if (leadModelConfig) {
setLeadModel(leadModelConfig as string);
setIsEnabled(true);
}
if (leadProviderConfig) setLeadProvider(leadProviderConfig as string);
if (leadTurnsConfig) setLeadTurns(Number(leadTurnsConfig));
if (failureThresholdConfig) setFailureThreshold(Number(failureThresholdConfig));
if (fallbackTurnsConfig) setFallbackTurns(Number(fallbackTurnsConfig));
// Set worker model to current model or from config
const workerModelConfig = await read('GOOSE_MODEL', false);
if (workerModelConfig) {
setWorkerModel(workerModelConfig as string);
} else if (currentModel) {
setWorkerModel(currentModel as string);
}
const workerProviderConfig = await read('GOOSE_PROVIDER', false);
if (workerProviderConfig) {
setWorkerProvider(workerProviderConfig as string);
}
// Load available models
const providers = await getProviders(false);
const activeProviders = providers.filter((p) => p.is_configured);
const options: { value: string; label: string; provider: string }[] = [];
activeProviders.forEach(({ metadata, name }) => {
if (metadata.known_models) {
metadata.known_models.forEach((model) => {
options.push({
value: model.name,
label: `${model.name} (${metadata.display_name})`,
provider: name,
});
});
}
});
setModelOptions(options);
} catch (error) {
console.error('Error loading configuration:', error);
} finally {
setIsLoading(false);
}
};
loadConfig();
}, [read, getProviders, currentModel]);
const handleSave = async () => {
try {
if (isEnabled && leadModel && workerModel) {
// Save lead/worker configuration
await Promise.all([
upsert('GOOSE_LEAD_MODEL', leadModel, false),
leadProvider && upsert('GOOSE_LEAD_PROVIDER', leadProvider, false),
upsert('GOOSE_MODEL', workerModel, false),
workerProvider && upsert('GOOSE_PROVIDER', workerProvider, false),
upsert('GOOSE_LEAD_TURNS', leadTurns, false),
upsert('GOOSE_LEAD_FAILURE_THRESHOLD', failureThreshold, false),
upsert('GOOSE_LEAD_FALLBACK_TURNS', fallbackTurns, false),
]);
} else {
// Remove lead/worker configuration
await Promise.all([
remove('GOOSE_LEAD_MODEL', false),
remove('GOOSE_LEAD_PROVIDER', false),
remove('GOOSE_LEAD_TURNS', false),
remove('GOOSE_LEAD_FAILURE_THRESHOLD', false),
remove('GOOSE_LEAD_FALLBACK_TURNS', false),
]);
}
onClose();
} catch (error) {
console.error('Error saving configuration:', error);
}
};
if (isLoading) {
return <div className="p-4">Loading...</div>;
}
return (
<div className="p-4 space-y-4">
<div className="space-y-2">
<h3 className="text-lg font-medium text-textProminent">Lead/Worker Mode</h3>
<p className="text-sm text-textSubtle">
Configure a lead model for planning and a worker model for execution
</p>
</div>
<div className="flex items-center space-x-2">
<input
type="checkbox"
id="enable-lead-worker"
checked={isEnabled}
onChange={(e) => setIsEnabled(e.target.checked)}
className="rounded border-borderStandard"
/>
<label htmlFor="enable-lead-worker" className="text-sm text-textStandard">
Enable lead/worker mode
</label>
</div>
<div className="space-y-4">
<div className="space-y-2">
<label className="text-sm text-textSubtle">Lead Model</label>
<Select
options={modelOptions}
value={modelOptions.find((opt) => opt.value === leadModel) || null}
onChange={(newValue: unknown) => {
const option = newValue as { value: string; provider: string } | null;
if (option) {
setLeadModel(option.value);
setLeadProvider(option.provider);
}
}}
placeholder="Select lead model..."
isDisabled={!isEnabled}
/>
<p className="text-xs text-textSubtle">
Strong model for initial planning and fallback recovery
</p>
</div>
<div className="space-y-2">
<label className="text-sm text-textSubtle">Worker Model</label>
<Select
options={modelOptions}
value={modelOptions.find((opt) => opt.value === workerModel) || null}
onChange={(newValue: unknown) => {
const option = newValue as { value: string; provider: string } | null;
if (option) {
setWorkerModel(option.value);
setWorkerProvider(option.provider);
}
}}
placeholder="Select worker model..."
isDisabled={!isEnabled}
/>
<p className="text-xs text-textSubtle">
Fast model for routine execution tasks
</p>
</div>
<div className="space-y-4 pt-4 border-t border-borderSubtle">
<div className="space-y-2">
<label className="text-sm text-textSubtle flex items-center gap-1">
Initial Lead Turns
<Info size={14} className="text-textSubtle" />
</label>
<Input
type="number"
min={1}
max={10}
value={leadTurns}
onChange={(e) => setLeadTurns(Number(e.target.value))}
className="w-20"
disabled={!isEnabled}
/>
<p className="text-xs text-textSubtle">
Number of turns to use the lead model at the start
</p>
</div>
<div className="space-y-2">
<label className="text-sm text-textSubtle flex items-center gap-1">
Failure Threshold
<Info size={14} className="text-textSubtle" />
</label>
<Input
type="number"
min={1}
max={5}
value={failureThreshold}
onChange={(e) => setFailureThreshold(Number(e.target.value))}
className="w-20"
disabled={!isEnabled}
/>
<p className="text-xs text-textSubtle">
Consecutive failures before switching back to lead
</p>
</div>
<div className="space-y-2">
<label className="text-sm text-textSubtle flex items-center gap-1">
Fallback Turns
<Info size={14} className="text-textSubtle" />
</label>
<Input
type="number"
min={1}
max={5}
value={fallbackTurns}
onChange={(e) => setFallbackTurns(Number(e.target.value))}
className="w-20"
disabled={!isEnabled}
/>
<p className="text-xs text-textSubtle">
Turns to use lead model during fallback
</p>
</div>
</div>
</div>
<div className="flex justify-end space-x-2 pt-4 border-t border-borderSubtle">
<Button variant="ghost" onClick={onClose}>
Cancel
</Button>
<Button onClick={handleSave} disabled={isEnabled && (!leadModel || !workerModel)}>
Save Settings
</Button>
</div>
</div>
);
}
+10
View File
@@ -0,0 +1,10 @@
import { useCurrentModelInfo } from '../components/ChatView';
export function useCurrentModel() {
const modelInfo = useCurrentModelInfo();
return {
currentModel: modelInfo?.model || null,
isLoading: false
};
}
+16
View File
@@ -24,6 +24,7 @@ type MessageEvent =
| { type: 'Message'; message: Message }
| { type: 'Error'; error: string }
| { type: 'Finish'; reason: string }
| { type: 'ModelChange'; model: string; mode: string }
| NotificationEvent;
export interface UseMessageStreamOptions {
@@ -140,6 +141,9 @@ export interface UseMessageStreamHelpers {
updateMessageStreamBody?: (newBody: object) => void;
notifications: NotificationEvent[];
/** Current model info from the backend */
currentModelInfo: { model: string; mode: string } | null;
}
/**
@@ -168,6 +172,7 @@ export function useMessageStream({
});
const [notifications, setNotifications] = useState<NotificationEvent[]>([]);
const [currentModelInfo, setCurrentModelInfo] = useState<{ model: string; mode: string } | null>(null);
// expose a way to update the body so we can update the session id when CLE occurs
const updateMessageStreamBody = useCallback((newBody: object) => {
@@ -273,6 +278,16 @@ export function useMessageStream({
break;
}
case 'ModelChange': {
// Update the current model in the frontend
const modelInfo = {
model: parsedEvent.model,
mode: parsedEvent.mode,
};
setCurrentModelInfo(modelInfo);
break;
}
case 'Error':
throw new Error(parsedEvent.error);
@@ -543,5 +558,6 @@ export function useMessageStream({
addToolResult,
updateMessageStreamBody,
notifications,
currentModelInfo,
};
}