Spaces:
Sleeping
Sleeping
NeonCharlie-24 commited on
Feat/vllm endpoints (#36)
Browse files* added initial tests for vllm client class.
* added the ImprovedVllmClient class.
* added vLLM option to config, bootstrap, and api routes.
* added vLLM option to front-end provider dropdown.
* added api key to docker-compose file and removed it from the yaml config files.
* added error handling and design decision comments.
* removed example vllm client file.
* added model re-discovery on request failure.
* convert vLLM to a single client and format persona roles to be vLLM acceptable names (system/user/assistant).
* moved _clean_responses into LLMClient base class to eliminate duplication.
* changed vLLM API label color to red and also set collapsed button color from grey to the respective providers color.
- docker-compose.yml +1 -0
- multi_llm_chatbot_backend/app/api/routes/provider.py +10 -0
- multi_llm_chatbot_backend/app/config.py +6 -0
- multi_llm_chatbot_backend/app/core/bootstrap.py +9 -1
- multi_llm_chatbot_backend/app/core/context_manager.py +21 -1
- multi_llm_chatbot_backend/app/llm/improved_gemini_client.py +0 -9
- multi_llm_chatbot_backend/app/llm/improved_vllm_client.py +67 -0
- multi_llm_chatbot_backend/app/llm/llm_client.py +9 -1
- multi_llm_chatbot_backend/app/tests/unit/test_vllm_client.py +175 -0
- multi_llm_chatbot_backend/requirements.txt +1 -0
- phd-advisor-frontend/src/components/ProviderDropdown.js +9 -2
- phd-advisor-frontend/src/styles/ChatPage.css +9 -0
- phd_config.yaml +2 -0
- undergrad_config.yaml +2 -0
docker-compose.yml
CHANGED
|
@@ -15,6 +15,7 @@ services:
|
|
| 15 |
MONGODB_DATABASE: phd_advisor
|
| 16 |
JWT_SECRET_KEY: ${JWT_SECRET_KEY:-CHANGEME-by-overriding-in-dot-env-file}
|
| 17 |
GEMINI_API_KEY: ${GEMINI_API_KEY:-?}
|
|
|
|
| 18 |
CORS_ORIGINS: ${CORS_ORIGINS:-http://localhost:3000}
|
| 19 |
GEMINI_MODEL: gemini-2.5-flash
|
| 20 |
CONFIG_PATH: ${CONFIG_PATH:-/ccai/phd_config.yaml}
|
|
|
|
| 15 |
MONGODB_DATABASE: phd_advisor
|
| 16 |
JWT_SECRET_KEY: ${JWT_SECRET_KEY:-CHANGEME-by-overriding-in-dot-env-file}
|
| 17 |
GEMINI_API_KEY: ${GEMINI_API_KEY:-?}
|
| 18 |
+
VLLM_API_KEY: ${VLLM_API_KEY:-}
|
| 19 |
CORS_ORIGINS: ${CORS_ORIGINS:-http://localhost:3000}
|
| 20 |
GEMINI_MODEL: gemini-2.5-flash
|
| 21 |
CONFIG_PATH: ${CONFIG_PATH:-/ccai/phd_config.yaml}
|
multi_llm_chatbot_backend/app/api/routes/provider.py
CHANGED
|
@@ -1,6 +1,8 @@
|
|
| 1 |
from fastapi import APIRouter, Body, HTTPException
|
|
|
|
| 2 |
from app.llm.improved_gemini_client import ImprovedGeminiClient
|
| 3 |
from app.llm.improved_ollama_client import ImprovedOllamaClient
|
|
|
|
| 4 |
from app.models.default_personas import get_default_personas
|
| 5 |
from app.core.bootstrap import chat_orchestrator, llm, current_provider, available_providers
|
| 6 |
from pydantic import BaseModel
|
|
@@ -24,6 +26,14 @@ def create_llm_client(provider: str = None):
|
|
| 24 |
return ImprovedOllamaClient(model_name="llama3.2:1b")
|
| 25 |
elif provider == "ollama":
|
| 26 |
return ImprovedOllamaClient(model_name="llama3.2:1b")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 27 |
else:
|
| 28 |
raise ValueError(f"Unknown provider: {provider}")
|
| 29 |
|
|
|
|
| 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 pydantic import BaseModel
|
|
|
|
| 26 |
return ImprovedOllamaClient(model_name="llama3.2:1b")
|
| 27 |
elif provider == "ollama":
|
| 28 |
return ImprovedOllamaClient(model_name="llama3.2:1b")
|
| 29 |
+
elif provider == "vllm":
|
| 30 |
+
settings = get_settings()
|
| 31 |
+
if not settings.llm.vllm.api_url:
|
| 32 |
+
raise ValueError("No vLLM endpoint configured. Set llm.vllm.api_url in your config.")
|
| 33 |
+
return ImprovedVllmClient(
|
| 34 |
+
api_url=settings.llm.vllm.api_url,
|
| 35 |
+
api_key=settings.llm.vllm.api_key,
|
| 36 |
+
)
|
| 37 |
else:
|
| 38 |
raise ValueError(f"Unknown provider: {provider}")
|
| 39 |
|
multi_llm_chatbot_backend/app/config.py
CHANGED
|
@@ -218,9 +218,15 @@ class OllamaConfig(BaseModel):
|
|
| 218 |
base_url: str = Field(default=os.getenv("OLLAMA_BASE_URL", "http://localhost:11434"))
|
| 219 |
|
| 220 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 221 |
class LLMConfig(BaseModel):
|
| 222 |
gemini: GeminiConfig = GeminiConfig()
|
| 223 |
ollama: OllamaConfig = OllamaConfig()
|
|
|
|
| 224 |
|
| 225 |
|
| 226 |
class RAGConfig(BaseModel):
|
|
|
|
| 218 |
base_url: str = Field(default=os.getenv("OLLAMA_BASE_URL", "http://localhost:11434"))
|
| 219 |
|
| 220 |
|
| 221 |
+
class VllmConfig(BaseModel):
|
| 222 |
+
api_url: str = ""
|
| 223 |
+
api_key: str = Field(default=os.getenv("VLLM_API_KEY", ""))
|
| 224 |
+
|
| 225 |
+
|
| 226 |
class LLMConfig(BaseModel):
|
| 227 |
gemini: GeminiConfig = GeminiConfig()
|
| 228 |
ollama: OllamaConfig = OllamaConfig()
|
| 229 |
+
vllm: VllmConfig = VllmConfig()
|
| 230 |
|
| 231 |
|
| 232 |
class RAGConfig(BaseModel):
|
multi_llm_chatbot_backend/app/core/bootstrap.py
CHANGED
|
@@ -2,19 +2,27 @@
|
|
| 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.core.improved_orchestrator import ImprovedChatOrchestrator
|
| 6 |
from app.models.default_personas import get_default_personas
|
| 7 |
|
| 8 |
settings = get_settings()
|
| 9 |
|
| 10 |
current_provider = "gemini"
|
| 11 |
-
available_providers = ["ollama", "gemini"]
|
| 12 |
|
| 13 |
def create_llm_client(provider=None):
|
| 14 |
if provider is None:
|
| 15 |
provider = current_provider
|
| 16 |
if provider == "gemini":
|
| 17 |
return ImprovedGeminiClient(model_name=settings.llm.gemini.model)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 18 |
else:
|
| 19 |
return ImprovedOllamaClient(
|
| 20 |
model_name=settings.llm.ollama.model,
|
|
|
|
| 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.core.improved_orchestrator import ImprovedChatOrchestrator
|
| 7 |
from app.models.default_personas import get_default_personas
|
| 8 |
|
| 9 |
settings = get_settings()
|
| 10 |
|
| 11 |
current_provider = "gemini"
|
| 12 |
+
available_providers = ["ollama", "gemini", "vllm"]
|
| 13 |
|
| 14 |
def create_llm_client(provider=None):
|
| 15 |
if provider is None:
|
| 16 |
provider = current_provider
|
| 17 |
if provider == "gemini":
|
| 18 |
return ImprovedGeminiClient(model_name=settings.llm.gemini.model)
|
| 19 |
+
elif provider == "vllm":
|
| 20 |
+
if not settings.llm.vllm.api_url:
|
| 21 |
+
raise ValueError("No vLLM endpoint configured. Set llm.vllm.api_url in your config.")
|
| 22 |
+
return ImprovedVllmClient(
|
| 23 |
+
api_url=settings.llm.vllm.api_url,
|
| 24 |
+
api_key=settings.llm.vllm.api_key,
|
| 25 |
+
)
|
| 26 |
else:
|
| 27 |
return ImprovedOllamaClient(
|
| 28 |
model_name=settings.llm.ollama.model,
|
multi_llm_chatbot_backend/app/core/context_manager.py
CHANGED
|
@@ -134,10 +134,30 @@ class ContextManager:
|
|
| 134 |
return self._format_for_gemini(messages, system_prompt)
|
| 135 |
elif provider.lower() in ["ollama", "mistral"]:
|
| 136 |
return self._format_for_ollama(messages, system_prompt)
|
|
|
|
|
|
|
| 137 |
else:
|
| 138 |
# Default format
|
| 139 |
return [{"role": "system", "content": system_prompt}] + messages
|
| 140 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 141 |
def _format_for_gemini(self, messages: List[dict], system_prompt: str) -> List[dict]:
|
| 142 |
"""
|
| 143 |
Format messages for Gemini API (uses user/model roles with parts structure)
|
|
|
|
| 134 |
return self._format_for_gemini(messages, system_prompt)
|
| 135 |
elif provider.lower() in ["ollama", "mistral"]:
|
| 136 |
return self._format_for_ollama(messages, system_prompt)
|
| 137 |
+
elif provider.lower() == "vllm":
|
| 138 |
+
return self._format_for_vllm(messages, system_prompt)
|
| 139 |
else:
|
| 140 |
# Default format
|
| 141 |
return [{"role": "system", "content": system_prompt}] + messages
|
| 142 |
+
|
| 143 |
+
def _format_for_vllm(self, messages: List[dict], system_prompt: str) -> List[dict]:
|
| 144 |
+
"""
|
| 145 |
+
Format messages for vLLM's OpenAI-compatible API.
|
| 146 |
+
Normalizes custom persona roles to 'assistant' since the API
|
| 147 |
+
only accepts system/user/assistant.
|
| 148 |
+
"""
|
| 149 |
+
formatted = [{"role": "system", "content": system_prompt}]
|
| 150 |
+
for message in messages:
|
| 151 |
+
role = message["role"]
|
| 152 |
+
content = message["content"]
|
| 153 |
+
if role == "system":
|
| 154 |
+
continue
|
| 155 |
+
if role not in ("user", "assistant"):
|
| 156 |
+
content = f"[{role.title()} Advisor]: {content}"
|
| 157 |
+
role = "assistant"
|
| 158 |
+
formatted.append({"role": role, "content": content})
|
| 159 |
+
return formatted
|
| 160 |
+
|
| 161 |
def _format_for_gemini(self, messages: List[dict], system_prompt: str) -> List[dict]:
|
| 162 |
"""
|
| 163 |
Format messages for Gemini API (uses user/model roles with parts structure)
|
multi_llm_chatbot_backend/app/llm/improved_gemini_client.py
CHANGED
|
@@ -1,6 +1,5 @@
|
|
| 1 |
import httpx
|
| 2 |
import os
|
| 3 |
-
import re
|
| 4 |
from typing import List
|
| 5 |
from app.llm.llm_client import LLMClient
|
| 6 |
from app.core.context_manager import get_context_manager
|
|
@@ -121,11 +120,3 @@ class ImprovedGeminiClient(LLMClient):
|
|
| 121 |
except Exception as e:
|
| 122 |
logger.error(f"Unexpected error in Gemini client: {str(e)}")
|
| 123 |
return "I encountered an unexpected error. Please try again."
|
| 124 |
-
|
| 125 |
-
def _clean_response(self, response: str) -> str:
|
| 126 |
-
"""Clean up response text, preserving Markdown formatting."""
|
| 127 |
-
response = response.replace("\r\n", "\n").replace("\r", "\n")
|
| 128 |
-
lines = [ln.rstrip() for ln in response.split("\n")]
|
| 129 |
-
response = re.sub(r"\n{3,}", "\n\n", "\n".join(lines)).strip()
|
| 130 |
-
|
| 131 |
-
return response
|
|
|
|
| 1 |
import httpx
|
| 2 |
import os
|
|
|
|
| 3 |
from typing import List
|
| 4 |
from app.llm.llm_client import LLMClient
|
| 5 |
from app.core.context_manager import get_context_manager
|
|
|
|
| 120 |
except Exception as e:
|
| 121 |
logger.error(f"Unexpected error in Gemini client: {str(e)}")
|
| 122 |
return "I encountered an unexpected error. Please try again."
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
multi_llm_chatbot_backend/app/llm/improved_vllm_client.py
ADDED
|
@@ -0,0 +1,67 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from typing import List
|
| 2 |
+
from openai import AsyncOpenAI, APIConnectionError, APIStatusError
|
| 3 |
+
from app.llm.llm_client import LLMClient
|
| 4 |
+
from app.core.context_manager import get_context_manager
|
| 5 |
+
import logging
|
| 6 |
+
|
| 7 |
+
logger = logging.getLogger(__name__)
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
class ImprovedVllmClient(LLMClient):
|
| 11 |
+
def __init__(self, api_url: str, api_key: str, model_name: str = None):
|
| 12 |
+
self.api_url = api_url
|
| 13 |
+
self.api_key = api_key
|
| 14 |
+
self.model_name = model_name
|
| 15 |
+
self.client = AsyncOpenAI(
|
| 16 |
+
base_url=f"{api_url}/v1",
|
| 17 |
+
api_key=api_key,
|
| 18 |
+
timeout=30.0,
|
| 19 |
+
)
|
| 20 |
+
self.context_manager = get_context_manager()
|
| 21 |
+
|
| 22 |
+
async def refresh_model(self):
|
| 23 |
+
"""Query the vLLM endpoint to discover the currently loaded model."""
|
| 24 |
+
models = await self.client.models.list()
|
| 25 |
+
if not models.data:
|
| 26 |
+
raise ValueError("No models available at the vLLM endpoint")
|
| 27 |
+
self.model_name = models.data[0].id
|
| 28 |
+
|
| 29 |
+
async def generate(self, system_prompt: str, context: List[dict],
|
| 30 |
+
temperature: float, max_tokens: int,
|
| 31 |
+
response_mime_type: str = None) -> str:
|
| 32 |
+
try:
|
| 33 |
+
context_window = self.context_manager.prepare_context_for_llm(
|
| 34 |
+
messages=context,
|
| 35 |
+
system_prompt=system_prompt,
|
| 36 |
+
llm_provider="vllm"
|
| 37 |
+
)
|
| 38 |
+
|
| 39 |
+
logger.debug(f"Context prepared: {len(context_window.messages)} messages, "
|
| 40 |
+
f"~{context_window.total_tokens} tokens, truncated={context_window.truncated}")
|
| 41 |
+
|
| 42 |
+
if not self.model_name:
|
| 43 |
+
await self.refresh_model()
|
| 44 |
+
|
| 45 |
+
response = await self.client.chat.completions.create(
|
| 46 |
+
model=self.model_name,
|
| 47 |
+
messages=context_window.messages,
|
| 48 |
+
temperature=temperature,
|
| 49 |
+
max_tokens=max_tokens,
|
| 50 |
+
)
|
| 51 |
+
|
| 52 |
+
text = response.choices[0].message.content.strip()
|
| 53 |
+
return self._clean_response(text)
|
| 54 |
+
|
| 55 |
+
except APIConnectionError:
|
| 56 |
+
logger.error(f"Unable to connect to vLLM at {self.api_url}")
|
| 57 |
+
return "I'm unable to connect to the AI service. Please ensure the vLLM endpoint is available."
|
| 58 |
+
except APIStatusError as e:
|
| 59 |
+
logger.error(f"vLLM API error: {e.status_code} - {e.message}")
|
| 60 |
+
if e.status_code == 404:
|
| 61 |
+
logger.info("Model not found, will re-discover on next request")
|
| 62 |
+
self.model_name = None
|
| 63 |
+
return "The AI service encountered an error. Please try again."
|
| 64 |
+
except Exception as e:
|
| 65 |
+
logger.error(f"Unexpected error in vLLM client: {str(e)}")
|
| 66 |
+
return "I encountered an unexpected error. Please try again."
|
| 67 |
+
|
multi_llm_chatbot_backend/app/llm/llm_client.py
CHANGED
|
@@ -1,5 +1,6 @@
|
|
| 1 |
from abc import ABC, abstractmethod
|
| 2 |
from typing import List
|
|
|
|
| 3 |
|
| 4 |
class LLMClient(ABC):
|
| 5 |
"""Abstract base class for all LLM clients"""
|
|
@@ -19,4 +20,11 @@ class LLMClient(ABC):
|
|
| 19 |
Returns:
|
| 20 |
str: The generated response text
|
| 21 |
"""
|
| 22 |
-
pass
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
from abc import ABC, abstractmethod
|
| 2 |
from typing import List
|
| 3 |
+
import re
|
| 4 |
|
| 5 |
class LLMClient(ABC):
|
| 6 |
"""Abstract base class for all LLM clients"""
|
|
|
|
| 20 |
Returns:
|
| 21 |
str: The generated response text
|
| 22 |
"""
|
| 23 |
+
pass
|
| 24 |
+
|
| 25 |
+
def _clean_response(self, response: str) -> str:
|
| 26 |
+
"""Clean up response text, preserving Markdown formatting."""
|
| 27 |
+
response = response.replace("\r\n", "\n").replace("\r", "\n")
|
| 28 |
+
lines = [ln.rstrip() for ln in response.split("\n")]
|
| 29 |
+
response = re.sub(r"\n{3,}", "\n\n", "\n".join(lines)).strip()
|
| 30 |
+
return response
|
multi_llm_chatbot_backend/app/tests/unit/test_vllm_client.py
ADDED
|
@@ -0,0 +1,175 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import asyncio
|
| 2 |
+
import unittest
|
| 3 |
+
from unittest.mock import AsyncMock, MagicMock, patch
|
| 4 |
+
|
| 5 |
+
from openai import APIConnectionError, APIStatusError
|
| 6 |
+
|
| 7 |
+
from app.llm.improved_vllm_client import ImprovedVllmClient
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
FAKE_URL = "https://fake.example.com/vllm0"
|
| 11 |
+
FAKE_KEY = "test-key"
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
def _make_completion_mock(content="Response"):
|
| 15 |
+
"""Build a mock that looks like an OpenAI ChatCompletion."""
|
| 16 |
+
mock_message = MagicMock()
|
| 17 |
+
mock_message.content = content
|
| 18 |
+
mock_choice = MagicMock()
|
| 19 |
+
mock_choice.message = mock_message
|
| 20 |
+
return MagicMock(choices=[mock_choice])
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
@patch("app.llm.improved_vllm_client.get_context_manager")
|
| 24 |
+
@patch("app.llm.improved_vllm_client.AsyncOpenAI")
|
| 25 |
+
class TestImprovedVllmClient(unittest.TestCase):
|
| 26 |
+
|
| 27 |
+
# ------------------------------------------------------------------
|
| 28 |
+
# Construction
|
| 29 |
+
# ------------------------------------------------------------------
|
| 30 |
+
|
| 31 |
+
def test_constructor_stores_attributes(self, MockAsyncOpenAI, mock_get_ctx):
|
| 32 |
+
client = ImprovedVllmClient(
|
| 33 |
+
api_url=FAKE_URL, api_key=FAKE_KEY, model_name="test-model",
|
| 34 |
+
)
|
| 35 |
+
self.assertEqual(client.api_url, FAKE_URL)
|
| 36 |
+
self.assertEqual(client.api_key, FAKE_KEY)
|
| 37 |
+
self.assertEqual(client.model_name, "test-model")
|
| 38 |
+
|
| 39 |
+
def test_constructor_defaults_model_to_none(self, MockAsyncOpenAI, mock_get_ctx):
|
| 40 |
+
client = ImprovedVllmClient(api_url=FAKE_URL, api_key=FAKE_KEY)
|
| 41 |
+
self.assertIsNone(client.model_name)
|
| 42 |
+
|
| 43 |
+
# ------------------------------------------------------------------
|
| 44 |
+
# Model discovery
|
| 45 |
+
# ------------------------------------------------------------------
|
| 46 |
+
|
| 47 |
+
def test_refresh_model_discovers_model(self, MockAsyncOpenAI, mock_get_ctx):
|
| 48 |
+
client = ImprovedVllmClient(api_url=FAKE_URL, api_key=FAKE_KEY)
|
| 49 |
+
|
| 50 |
+
mock_model = MagicMock()
|
| 51 |
+
mock_model.id = "discovered-model"
|
| 52 |
+
client.client.models.list = AsyncMock(
|
| 53 |
+
return_value=MagicMock(data=[mock_model])
|
| 54 |
+
)
|
| 55 |
+
|
| 56 |
+
asyncio.run(client.refresh_model())
|
| 57 |
+
self.assertEqual(client.model_name, "discovered-model")
|
| 58 |
+
|
| 59 |
+
# ------------------------------------------------------------------
|
| 60 |
+
# generate – happy path
|
| 61 |
+
# ------------------------------------------------------------------
|
| 62 |
+
|
| 63 |
+
def test_generate_returns_cleaned_response(self, MockAsyncOpenAI, mock_get_ctx):
|
| 64 |
+
client = ImprovedVllmClient(
|
| 65 |
+
api_url=FAKE_URL, api_key=FAKE_KEY, model_name="test-model",
|
| 66 |
+
)
|
| 67 |
+
client.client.chat.completions.create = AsyncMock(
|
| 68 |
+
return_value=_make_completion_mock(" Here is my response. ")
|
| 69 |
+
)
|
| 70 |
+
|
| 71 |
+
result = asyncio.run(client.generate(
|
| 72 |
+
system_prompt="You are helpful.",
|
| 73 |
+
context=[{"role": "user", "content": "Hello"}],
|
| 74 |
+
temperature=0.7,
|
| 75 |
+
max_tokens=100,
|
| 76 |
+
))
|
| 77 |
+
self.assertEqual(result, "Here is my response.")
|
| 78 |
+
|
| 79 |
+
def test_generate_auto_discovers_model_when_none(self, MockAsyncOpenAI, mock_get_ctx):
|
| 80 |
+
client = ImprovedVllmClient(
|
| 81 |
+
api_url=FAKE_URL, api_key=FAKE_KEY, model_name=None,
|
| 82 |
+
)
|
| 83 |
+
|
| 84 |
+
mock_model = MagicMock()
|
| 85 |
+
mock_model.id = "auto-discovered"
|
| 86 |
+
client.client.models.list = AsyncMock(
|
| 87 |
+
return_value=MagicMock(data=[mock_model])
|
| 88 |
+
)
|
| 89 |
+
client.client.chat.completions.create = AsyncMock(
|
| 90 |
+
return_value=_make_completion_mock()
|
| 91 |
+
)
|
| 92 |
+
|
| 93 |
+
asyncio.run(client.generate(
|
| 94 |
+
system_prompt="Test",
|
| 95 |
+
context=[{"role": "user", "content": "Hi"}],
|
| 96 |
+
temperature=0.5,
|
| 97 |
+
max_tokens=50,
|
| 98 |
+
))
|
| 99 |
+
|
| 100 |
+
client.client.models.list.assert_called_once()
|
| 101 |
+
self.assertEqual(client.model_name, "auto-discovered")
|
| 102 |
+
|
| 103 |
+
# ------------------------------------------------------------------
|
| 104 |
+
# generate – error handling
|
| 105 |
+
# ------------------------------------------------------------------
|
| 106 |
+
|
| 107 |
+
def test_generate_handles_connection_error(self, MockAsyncOpenAI, mock_get_ctx):
|
| 108 |
+
client = ImprovedVllmClient(
|
| 109 |
+
api_url=FAKE_URL, api_key=FAKE_KEY, model_name="test-model",
|
| 110 |
+
)
|
| 111 |
+
client.client.chat.completions.create = AsyncMock(
|
| 112 |
+
side_effect=APIConnectionError(request=MagicMock())
|
| 113 |
+
)
|
| 114 |
+
|
| 115 |
+
result = asyncio.run(client.generate(
|
| 116 |
+
system_prompt="Test",
|
| 117 |
+
context=[{"role": "user", "content": "Hi"}],
|
| 118 |
+
temperature=0.5,
|
| 119 |
+
max_tokens=50,
|
| 120 |
+
))
|
| 121 |
+
self.assertIn("unable to connect", result.lower())
|
| 122 |
+
|
| 123 |
+
def test_generate_handles_status_error(self, MockAsyncOpenAI, mock_get_ctx):
|
| 124 |
+
client = ImprovedVllmClient(
|
| 125 |
+
api_url=FAKE_URL, api_key=FAKE_KEY, model_name="test-model",
|
| 126 |
+
)
|
| 127 |
+
mock_response = MagicMock()
|
| 128 |
+
mock_response.status_code = 500
|
| 129 |
+
client.client.chat.completions.create = AsyncMock(
|
| 130 |
+
side_effect=APIStatusError(
|
| 131 |
+
message="Server error", response=mock_response, body=None,
|
| 132 |
+
)
|
| 133 |
+
)
|
| 134 |
+
|
| 135 |
+
result = asyncio.run(client.generate(
|
| 136 |
+
system_prompt="Test",
|
| 137 |
+
context=[{"role": "user", "content": "Hi"}],
|
| 138 |
+
temperature=0.5,
|
| 139 |
+
max_tokens=50,
|
| 140 |
+
))
|
| 141 |
+
self.assertIn("error", result.lower())
|
| 142 |
+
|
| 143 |
+
def test_generate_clears_model_on_404(self, MockAsyncOpenAI, mock_get_ctx):
|
| 144 |
+
client = ImprovedVllmClient(
|
| 145 |
+
api_url=FAKE_URL, api_key=FAKE_KEY, model_name="stale-model",
|
| 146 |
+
)
|
| 147 |
+
mock_response = MagicMock()
|
| 148 |
+
mock_response.status_code = 404
|
| 149 |
+
client.client.chat.completions.create = AsyncMock(
|
| 150 |
+
side_effect=APIStatusError(
|
| 151 |
+
message="Model not found", response=mock_response, body=None,
|
| 152 |
+
)
|
| 153 |
+
)
|
| 154 |
+
|
| 155 |
+
asyncio.run(client.generate(
|
| 156 |
+
system_prompt="Test",
|
| 157 |
+
context=[{"role": "user", "content": "Hi"}],
|
| 158 |
+
temperature=0.5,
|
| 159 |
+
max_tokens=50,
|
| 160 |
+
))
|
| 161 |
+
self.assertIsNone(client.model_name)
|
| 162 |
+
|
| 163 |
+
# ------------------------------------------------------------------
|
| 164 |
+
# _clean_response
|
| 165 |
+
# ------------------------------------------------------------------
|
| 166 |
+
|
| 167 |
+
def test_clean_response_normalizes_whitespace(self, MockAsyncOpenAI, mock_get_ctx):
|
| 168 |
+
client = ImprovedVllmClient(
|
| 169 |
+
api_url=FAKE_URL, api_key=FAKE_KEY, model_name="test-model",
|
| 170 |
+
)
|
| 171 |
+
dirty = "Line one.\r\n\r\n\r\n\r\nLine two. "
|
| 172 |
+
cleaned = client._clean_response(dirty)
|
| 173 |
+
self.assertNotIn("\r", cleaned)
|
| 174 |
+
self.assertNotIn("\n\n\n", cleaned)
|
| 175 |
+
self.assertEqual(cleaned, "Line one.\n\nLine two.")
|
multi_llm_chatbot_backend/requirements.txt
CHANGED
|
@@ -5,6 +5,7 @@ python-multipart
|
|
| 5 |
|
| 6 |
# HTTP client for LLM APIs
|
| 7 |
httpx
|
|
|
|
| 8 |
|
| 9 |
# Document processing
|
| 10 |
PyPDF2
|
|
|
|
| 5 |
|
| 6 |
# HTTP client for LLM APIs
|
| 7 |
httpx
|
| 8 |
+
openai~=2.30
|
| 9 |
|
| 10 |
# Document processing
|
| 11 |
PyPDF2
|
phd-advisor-frontend/src/components/ProviderDropdown.js
CHANGED
|
@@ -1,6 +1,6 @@
|
|
| 1 |
// src/components/ProviderDropdown.js
|
| 2 |
import React, { useState, useRef, useEffect } from 'react';
|
| 3 |
-
import { ChevronDown, Cpu, Cloud, Loader2 } from 'lucide-react';
|
| 4 |
import { useTheme } from '../contexts/ThemeContext';
|
| 5 |
|
| 6 |
const ProviderDropdown = ({ currentProvider, onProviderChange, isLoading = false }) => {
|
|
@@ -22,6 +22,13 @@ const ProviderDropdown = ({ currentProvider, onProviderChange, isLoading = false
|
|
| 22 |
description: 'Local LLM via Ollama',
|
| 23 |
icon: Cpu,
|
| 24 |
badge: 'Local'
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 25 |
}
|
| 26 |
];
|
| 27 |
|
|
@@ -69,7 +76,7 @@ const ProviderDropdown = ({ currentProvider, onProviderChange, isLoading = false
|
|
| 69 |
)}
|
| 70 |
<div className="provider-info">
|
| 71 |
<span className="provider-name">{currentProviderInfo?.name || 'Unknown'}</span>
|
| 72 |
-
<span className=
|
| 73 |
</div>
|
| 74 |
</div>
|
| 75 |
<ChevronDown
|
|
|
|
| 1 |
// src/components/ProviderDropdown.js
|
| 2 |
import React, { useState, useRef, useEffect } from 'react';
|
| 3 |
+
import { ChevronDown, Cpu, Cloud, Server, Loader2 } from 'lucide-react';
|
| 4 |
import { useTheme } from '../contexts/ThemeContext';
|
| 5 |
|
| 6 |
const ProviderDropdown = ({ currentProvider, onProviderChange, isLoading = false }) => {
|
|
|
|
| 22 |
description: 'Local LLM via Ollama',
|
| 23 |
icon: Cpu,
|
| 24 |
badge: 'Local'
|
| 25 |
+
},
|
| 26 |
+
{
|
| 27 |
+
id: 'vllm',
|
| 28 |
+
name: 'vLLM',
|
| 29 |
+
description: 'vLLM inference endpoint',
|
| 30 |
+
icon: Server,
|
| 31 |
+
badge: 'API'
|
| 32 |
}
|
| 33 |
];
|
| 34 |
|
|
|
|
| 76 |
)}
|
| 77 |
<div className="provider-info">
|
| 78 |
<span className="provider-name">{currentProviderInfo?.name || 'Unknown'}</span>
|
| 79 |
+
<span className={`provider-badge ${currentProvider}`}>{currentProviderInfo?.badge}</span>
|
| 80 |
</div>
|
| 81 |
</div>
|
| 82 |
<ChevronDown
|
phd-advisor-frontend/src/styles/ChatPage.css
CHANGED
|
@@ -882,6 +882,10 @@
|
|
| 882 |
letter-spacing: 0.5px;
|
| 883 |
}
|
| 884 |
|
|
|
|
|
|
|
|
|
|
|
|
|
| 885 |
.clarification-message-container {
|
| 886 |
display: flex;
|
| 887 |
justify-content: center;
|
|
@@ -1112,6 +1116,11 @@
|
|
| 1112 |
color: #10b981;
|
| 1113 |
}
|
| 1114 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1115 |
.provider-option-description {
|
| 1116 |
font-size: 11px;
|
| 1117 |
color: var(--text-tertiary);
|
|
|
|
| 882 |
letter-spacing: 0.5px;
|
| 883 |
}
|
| 884 |
|
| 885 |
+
.provider-badge.gemini { color: #4285f4; }
|
| 886 |
+
.provider-badge.ollama { color: #10b981; }
|
| 887 |
+
.provider-badge.vllm { color: #ef4444; }
|
| 888 |
+
|
| 889 |
.clarification-message-container {
|
| 890 |
display: flex;
|
| 891 |
justify-content: center;
|
|
|
|
| 1116 |
color: #10b981;
|
| 1117 |
}
|
| 1118 |
|
| 1119 |
+
.provider-option-badge.vllm {
|
| 1120 |
+
background: rgba(239, 68, 68, 0.15);
|
| 1121 |
+
color: #ef4444;
|
| 1122 |
+
}
|
| 1123 |
+
|
| 1124 |
.provider-option-description {
|
| 1125 |
font-size: 11px;
|
| 1126 |
color: var(--text-tertiary);
|
phd_config.yaml
CHANGED
|
@@ -144,6 +144,8 @@ llm:
|
|
| 144 |
model: "gemini-2.5-flash"
|
| 145 |
ollama:
|
| 146 |
model: "llama3.2:1b"
|
|
|
|
|
|
|
| 147 |
|
| 148 |
rag:
|
| 149 |
embedding_model: "all-MiniLM-L6-v2"
|
|
|
|
| 144 |
model: "gemini-2.5-flash"
|
| 145 |
ollama:
|
| 146 |
model: "llama3.2:1b"
|
| 147 |
+
vllm:
|
| 148 |
+
api_url: https://rtx6000blackwell-1.neonaiservices2.com/vllm0
|
| 149 |
|
| 150 |
rag:
|
| 151 |
embedding_model: "all-MiniLM-L6-v2"
|
undergrad_config.yaml
CHANGED
|
@@ -416,6 +416,8 @@ llm:
|
|
| 416 |
model: "gemini-2.5-flash"
|
| 417 |
ollama:
|
| 418 |
model: "llama3.2:1b"
|
|
|
|
|
|
|
| 419 |
|
| 420 |
rag:
|
| 421 |
embedding_model: "all-MiniLM-L6-v2"
|
|
|
|
| 416 |
model: "gemini-2.5-flash"
|
| 417 |
ollama:
|
| 418 |
model: "llama3.2:1b"
|
| 419 |
+
vllm:
|
| 420 |
+
api_url: https://rtx6000blackwell-1.neonaiservices2.com/vllm0
|
| 421 |
|
| 422 |
rag:
|
| 423 |
embedding_model: "all-MiniLM-L6-v2"
|