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:
@@ -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;
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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?;
|
||||
|
||||
@@ -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)]
|
||||
|
||||
@@ -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
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
@@ -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: {}",
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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>
|
||||
);
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
import { useCurrentModelInfo } from '../components/ChatView';
|
||||
|
||||
export function useCurrentModel() {
|
||||
const modelInfo = useCurrentModelInfo();
|
||||
|
||||
return {
|
||||
currentModel: modelInfo?.model || null,
|
||||
isLoading: false
|
||||
};
|
||||
}
|
||||
@@ -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,
|
||||
};
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user