Skip to content

Commit 300ed29

Browse files
Feat/support hybrid model selection rebase (#80)
* feat(hybrid): add Hybrid provider option to ProviderDropdown (UI scaffold) Adds a 4th "Hybrid" option to the provider dropdown with a purple "Mixed" badge and dark-mode fix, plus a Settings2 configure button on the selected hybrid row. Passes new `defaultBackend` and `backendLocked` (BrainForge) fields from `/api/config` personas through AppConfigContext for the upcoming mapping UI. Backend needs: - `default_backend` + `brainforge` flags on persona config - accept "hybrid" on POST /switch-provider - GET/POST /hybrid-config for { orchestrator, personas{} } map - per-persona inference routing; BrainForge hard-pinned * added per-user llm backend selection with uniform (all same provider) and hybird (different llm per persona/orchestrator) modes. * Enhance ChatPage with hybrid LLM configuration support - Added HybridConfigModal for configuring hybrid LLM settings. - Refactored state management to support both uniform and hybrid modes. -Each adviosr can now have its own model * added filter for avaliable backends with health status check set to ping every 300 seconds. * added unit tests for the available backends filter and health check. * Pending changes to build backend the menus will be moved to the welcome screens and settigns pages once merges are completed. * fixed bootstrap.py import causing test_available_backends failure. * added conftest.py to simplify mock module imports and unit tests for LLM provider config. * restored needs_clarification_improved function lost during rebase. * Enhance SettingsModal with user profile and account management features - Added functionality for updating user profile information (first name, last name). - Implemented password change and account deletion processes with confirmation. - Improved modal behavior to prevent accidental closure during text selection. - Updated ChatPage to integrate new SettingsModal features, including user update and sign-out callbacks. * fix stubbing issue in conftest.py from rebase. * fix backend config values to lock brainforge models on frontend and prevent their underlying models from being changed. * added admin-level enabled toggle to each provider and set default backend dynamically. * replaced frontend hardcoded gemini fallback with dynamic defaults. * fallback to default_backend when hybrid mode has no overrides. * add test case for gemini missing. * added configurable default_backend parameter to config.yaml. * Fixed the black on black text and added the default option --------- Co-authored-by: Neon:ryan <ryan@neon.ai>
1 parent ef41a67 commit 300ed29

25 files changed

Lines changed: 1587 additions & 217 deletions

‎multi_llm_chatbot_backend/app/api/routes/chat.py‎

Lines changed: 77 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,7 @@
1313
from app.api.utils import get_or_create_session_for_request_async
1414
from app.core.auth import get_current_active_user
1515
from app.config import get_settings
16-
from app.core.bootstrap import chat_orchestrator
16+
from app.core.bootstrap import chat_orchestrator, get_llm_client
1717
from app.core.database import get_database
1818
from app.core.persona_filter import get_available_persona_ids
1919
from app.core.session_manager import get_session_manager
@@ -24,6 +24,41 @@
2424
router = APIRouter()
2525
session_manager = get_session_manager()
2626

27+
28+
def resolve_llm_clients(user: User) -> Dict[str, Any]:
29+
"""Resolve LLM clients from a user's stored configuration.
30+
31+
Returns ``{"orchestrator": LLMClient | None, "personas": {id: LLMClient} | None}``.
32+
33+
- No saved config: both values are ``None``; callers fall back to
34+
orchestrator/persona defaults.
35+
- Uniform mode: the same cached client is returned for the orchestrator
36+
and every persona.
37+
- Hybrid mode: the orchestrator and each persona may receive different
38+
clients based on the user's per-persona mapping.
39+
"""
40+
config = user.llm_config
41+
if config is None:
42+
return {"orchestrator": None, "personas": None}
43+
44+
if config.mode == "uniform":
45+
client = get_llm_client(config.default_backend)
46+
persona_clients = {
47+
pid: client for pid in chat_orchestrator.personas
48+
}
49+
return {"orchestrator": client, "personas": persona_clients}
50+
51+
# Hybrid mode
52+
orchestrator_backend = config.orchestrator_backend or config.default_backend
53+
orchestrator_client = get_llm_client(orchestrator_backend)
54+
55+
persona_clients = {}
56+
for pid in chat_orchestrator.personas:
57+
backend = (config.persona_backends or {}).get(pid, config.default_backend)
58+
persona_clients[pid] = get_llm_client(backend)
59+
60+
return {"orchestrator": orchestrator_client, "personas": persona_clients}
61+
2762
# Enhanced data models
2863
class UserInput(BaseModel):
2964
user_input: str
@@ -81,6 +116,11 @@ async def chat_stream(
81116

82117
async def _event_generator():
83118
try:
119+
# Resolve per-user LLM clients from their stored config
120+
llm_clients = resolve_llm_clients(current_user)
121+
orchestrator_llm = llm_clients["orchestrator"]
122+
persona_llms = llm_clients["personas"]
123+
84124
# Load or create the in-memory session
85125
if message.chat_session_id:
86126
sid = f"chat_{message.chat_session_id}"
@@ -107,7 +147,9 @@ async def _event_generator():
107147
).to_ndjson()
108148

109149
if await chat_orchestrator.needs_clarification_improved(session, message.user_input):
110-
clar = await chat_orchestrator.generate_contextual_clarification(message.user_input)
150+
clar = await chat_orchestrator.generate_contextual_clarification(
151+
message.user_input, llm_client=orchestrator_llm,
152+
)
111153
yield ChatStreamLine(
112154
type="clarification",
113155
data={
@@ -123,7 +165,9 @@ async def _event_generator():
123165

124166
# If an enabled tool can handle this query, return its response
125167
# directly and skip persona generation.
126-
tool_result = await chat_orchestrator.get_tool_response(message.user_input)
168+
tool_result = await chat_orchestrator.get_tool_response(
169+
message.user_input, llm_client=orchestrator_llm,
170+
)
127171
if tool_result.used_tool:
128172
# Append user message to in-memory session and persist to MongoDB
129173
session.append_message("orchestrator", tool_result.text)
@@ -164,6 +208,7 @@ async def _event_generator():
164208
top_personas = await chat_orchestrator.get_top_personas(
165209
session_id=sid,
166210
allowed_ids=available,
211+
llm_client=orchestrator_llm,
167212
)
168213

169214
# Guard against race condition where all selected advisors
@@ -210,9 +255,11 @@ async def _run(pid: str) -> None:
210255
"document_chunks_used": 0,
211256
})
212257
return
258+
persona_llm = (persona_llms or {}).get(pid)
213259
result = await chat_orchestrator.generate_single_persona_response(
214260
session, persona,
215261
message.response_length or "medium",
262+
llm_client=persona_llm,
216263
)
217264
session.append_message(pid, result["response"])
218265
await done_queue.put(result)
@@ -390,7 +437,10 @@ async def create_new_chat(
390437
raise HTTPException(status_code=500, detail="Failed to create new chat")
391438

392439
@router.post("/chat/{persona_id}")
393-
async def chat_with_specific_advisor(persona_id: str, input: UserInput, request: Request):
440+
async def chat_with_specific_advisor(
441+
persona_id: str, input: UserInput, request: Request,
442+
current_user: User = Depends(get_current_active_user),
443+
):
394444
"""Chat with a specific advisor - UPDATED"""
395445
try:
396446
if persona_id not in chat_orchestrator.personas:
@@ -408,11 +458,15 @@ async def chat_with_specific_advisor(persona_id: str, input: UserInput, request:
408458
isExpandRequest=True,
409459
),
410460
)
461+
462+
llm_clients = resolve_llm_clients(current_user)
463+
persona_llm = (llm_clients["personas"] or {}).get(persona_id)
411464

412465
result = await chat_orchestrator.chat_with_persona(
413466
user_input=input.user_input,
414467
persona_id=persona_id,
415-
session_id=session_id
468+
session_id=session_id,
469+
llm_client=persona_llm,
416470
)
417471

418472
# Handle response structure
@@ -479,7 +533,10 @@ async def chat_with_specific_advisor(persona_id: str, input: UserInput, request:
479533
}
480534

481535
@router.post("/reply-to-advisor")
482-
async def reply_to_advisor(reply: ReplyToAdvisor, request: Request):
536+
async def reply_to_advisor(
537+
reply: ReplyToAdvisor, request: Request,
538+
current_user: User = Depends(get_current_active_user),
539+
):
483540
"""Reply to a specific advisor with proper context - UPDATED"""
484541
try:
485542
if reply.advisor_id not in chat_orchestrator.personas:
@@ -520,10 +577,14 @@ async def reply_to_advisor(reply: ReplyToAdvisor, request: Request):
520577
if original_message:
521578
contextual_input = f"[Replying to your previous message: '{original_message[:100]}...'] {reply.user_input}"
522579

580+
llm_clients = resolve_llm_clients(current_user)
581+
advisor_llm = (llm_clients["personas"] or {}).get(reply.advisor_id)
582+
523583
result = await chat_orchestrator.chat_with_persona(
524584
user_input=contextual_input,
525585
persona_id=reply.advisor_id,
526-
session_id=session_id
586+
session_id=session_id,
587+
llm_client=advisor_llm,
527588
)
528589

529590
# Handle response structure
@@ -600,15 +661,22 @@ async def reply_to_advisor(reply: ReplyToAdvisor, request: Request):
600661
}
601662

602663
@router.post("/ask/")
603-
async def ask_question(query: PersonaQuery, request: Request):
664+
async def ask_question(
665+
query: PersonaQuery, request: Request,
666+
current_user: User = Depends(get_current_active_user),
667+
):
604668
"""Ask question - UPDATED"""
605669
try:
606670
session_id = await get_or_create_session_for_request_async(request)
607671

672+
llm_clients = resolve_llm_clients(current_user)
673+
persona_llm = (llm_clients["personas"] or {}).get(query.persona)
674+
608675
result = await chat_orchestrator.chat_with_persona(
609676
user_input=query.question,
610677
persona_id=query.persona,
611-
session_id=session_id
678+
session_id=session_id,
679+
llm_client=persona_llm,
612680
)
613681

614682
if result["type"] == "single_persona_response":
Lines changed: 57 additions & 93 deletions
Original file line numberDiff line numberDiff line change
@@ -1,108 +1,72 @@
1-
from fastapi import APIRouter, Body, HTTPException
2-
from app.config import get_settings
3-
from app.llm.improved_gemini_client import ImprovedGeminiClient
4-
from app.llm.improved_ollama_client import ImprovedOllamaClient
5-
from app.llm.improved_vllm_client import ImprovedVllmClient
6-
from app.models.default_personas import get_default_personas
7-
from app.core.bootstrap import chat_orchestrator, llm, current_provider, available_providers
8-
from app.core.brainforge_sync import BRAINFORGE_PERSONA_PREFIX
9-
from pydantic import BaseModel
10-
import os
1+
from fastapi import APIRouter, Depends, HTTPException, status
2+
from app.core.auth import get_current_active_user
3+
from app.core.bootstrap import (
4+
chat_orchestrator, get_llm_client, AVAILABLE_BACKENDS, _is_backend_enabled,
5+
)
6+
from app.core.database import get_database
7+
from app.models.user import User, UserLLMConfig
118
import logging
129

1310
logger = logging.getLogger(__name__)
1411

1512
router = APIRouter()
1613

17-
def create_llm_client(provider: str = None):
18-
global current_provider
19-
if provider is None:
20-
provider = current_provider
21-
22-
if provider == "gemini":
23-
try:
24-
return ImprovedGeminiClient(model_name=os.getenv("GEMINI_MODEL"))
25-
except ValueError as e:
26-
logger.warning(f"Gemini API key not found, falling back to Ollama: {e}")
27-
return ImprovedOllamaClient(model_name="llama3.2:1b")
28-
elif provider == "ollama":
29-
return ImprovedOllamaClient(model_name="llama3.2:1b")
30-
elif provider == "vllm":
31-
settings = get_settings()
32-
if not settings.llm.vllm.api_url:
33-
raise ValueError("No vLLM endpoint configured. Set llm.vllm.api_url in your config.")
34-
return ImprovedVllmClient(
35-
api_url=settings.llm.vllm.api_url,
36-
api_key=settings.llm.vllm.api_key,
37-
)
38-
else:
39-
raise ValueError(f"Unknown provider: {provider}")
40-
41-
# Initialize LLM and personas
42-
llm = create_llm_client(current_provider)
43-
DEFAULT_PERSONAS = get_default_personas(llm)
44-
for persona in DEFAULT_PERSONAS:
45-
chat_orchestrator.register_persona(persona)
46-
47-
class ProviderSwitch(BaseModel):
48-
provider: str
4914

5015
@router.get("/current-provider")
51-
async def get_current_provider():
16+
async def get_current_provider(
17+
current_user: User = Depends(get_current_active_user),
18+
):
19+
"""Return the authenticated user's LLM configuration."""
20+
config = current_user.llm_config or UserLLMConfig()
5221
return {
53-
"current_provider": current_provider,
54-
"available_providers": available_providers,
55-
"model_info": {
56-
"name": llm.model_name if hasattr(llm, 'model_name') else "gemini-2.0-flash",
57-
"provider": current_provider
58-
}
22+
"llm_config": config.model_dump(),
23+
"available_backends": AVAILABLE_BACKENDS,
5924
}
6025

61-
@router.post("/switch-provider")
62-
async def switch_provider(provider_data: ProviderSwitch):
63-
global current_provider, llm
64-
65-
if provider_data.provider not in available_providers:
66-
raise HTTPException(status_code=400, detail=f"Unknown provider: {provider_data.provider}. Available: {available_providers}")
67-
68-
try:
69-
current_provider = provider_data.provider
70-
new_llm = create_llm_client(current_provider)
71-
llm = new_llm
7226

73-
chat_orchestrator.llm_client = new_llm
74-
75-
new_personas = get_default_personas(new_llm)
76-
# Clear only non-BrainForge personas; BF advisors have their own LLM clients
77-
non_bf_ids = [pid for pid in chat_orchestrator.personas if not pid.startswith(f"{BRAINFORGE_PERSONA_PREFIX}_")]
78-
for pid in non_bf_ids:
79-
chat_orchestrator.unregister_persona(pid)
80-
for persona in new_personas:
81-
chat_orchestrator.register_persona(persona)
82-
83-
return {
84-
"message": f"Successfully switched to {current_provider}",
85-
"current_provider": current_provider,
86-
"model_info": {
87-
"name": new_llm.model_name if hasattr(new_llm, 'model_name') else "gemini-2.0-flash",
88-
"provider": current_provider
89-
}
90-
}
91-
92-
except Exception as e:
93-
raise HTTPException(status_code=500, detail=f"Failed to switch to {provider_data.provider}: {str(e)}")
94-
95-
@router.post("/switch-model")
96-
async def switch_model(model_name: str = Body(...)):
97-
if "gemini" in model_name.lower():
98-
return await switch_provider(ProviderSwitch(provider="gemini"))
99-
else:
100-
return await switch_provider(ProviderSwitch(provider="ollama"))
27+
@router.post("/switch-provider")
28+
async def switch_provider(
29+
llm_config: UserLLMConfig,
30+
current_user: User = Depends(get_current_active_user),
31+
):
32+
"""Persist the user's LLM configuration to their profile."""
33+
if llm_config.mode == "hybrid" and llm_config.persona_backends:
34+
registered = set(chat_orchestrator.personas.keys())
35+
unknown = set(llm_config.persona_backends.keys()) - registered
36+
if unknown:
37+
raise HTTPException(
38+
status_code=status.HTTP_400_BAD_REQUEST,
39+
detail=f"Unknown persona IDs: {sorted(unknown)}. "
40+
f"Valid IDs: {sorted(registered)}",
41+
)
42+
43+
backends_to_check = {llm_config.default_backend}
44+
if llm_config.orchestrator_backend:
45+
backends_to_check.add(llm_config.orchestrator_backend)
46+
if llm_config.persona_backends:
47+
backends_to_check.update(llm_config.persona_backends.values())
48+
49+
for backend in backends_to_check:
50+
if not _is_backend_enabled(backend):
51+
raise HTTPException(
52+
status_code=status.HTTP_400_BAD_REQUEST,
53+
detail=f"Backend {backend!r} is disabled by the administrator.",
54+
)
55+
try:
56+
get_llm_client(backend)
57+
except Exception as exc:
58+
raise HTTPException(
59+
status_code=status.HTTP_400_BAD_REQUEST,
60+
detail=f"Backend {backend!r} is not configured: {exc}",
61+
)
62+
63+
db = get_database()
64+
await db.users.update_one(
65+
{"_id": current_user.id},
66+
{"$set": {"llm_config": llm_config.model_dump()}},
67+
)
10168

102-
@router.get("/current-model")
103-
async def get_current_model():
104-
model_name = llm.model_name if hasattr(llm, 'model_name') else "gemini-2.0-flash"
10569
return {
106-
"model": model_name,
107-
"provider": current_provider
70+
"message": "LLM configuration updated",
71+
"llm_config": llm_config.model_dump(),
10872
}

‎multi_llm_chatbot_backend/app/config.py‎

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -253,6 +253,7 @@ def _warn_connection_envvar(self):
253253

254254

255255
class GeminiConfig(BaseModel):
256+
enabled: bool = True
256257
api_key: str = Field(default=os.getenv("GEMINI_API_KEY"))
257258
model: str = "gemini-2.5-flash"
258259

@@ -272,12 +273,14 @@ def _warn_gemini_envvar(self):
272273

273274

274275
class OllamaConfig(BaseModel):
276+
enabled: bool = True
275277
model: str = "llama3.2:1b"
276278
# TODO: Drop support for `OLLAMA_BASE_URL` envvar handling
277279
base_url: str = Field(default=os.getenv("OLLAMA_BASE_URL", "http://localhost:11434"))
278280

279281

280282
class VllmConfig(BaseModel):
283+
enabled: bool = True
281284
api_url: str = ""
282285
api_key: str = Field(default=os.getenv("VLLM_API_KEY", ""))
283286

@@ -290,10 +293,12 @@ class BrainForgeConfig(BaseModel):
290293

291294

292295
class LLMConfig(BaseModel):
296+
default_backend: str = ""
293297
gemini: GeminiConfig = GeminiConfig()
294298
ollama: OllamaConfig = OllamaConfig()
295299
vllm: VllmConfig = VllmConfig()
296300
brainforge: BrainForgeConfig = BrainForgeConfig()
301+
health_check_interval_seconds: int = 300
297302

298303

299304
class RAGConfig(BaseModel):

0 commit comments

Comments
 (0)