Merge pull request !2834 from rhodecode-enterprise-ce feature/RCCE-325_AI-model-versions-auto-update
feature: implements possibility to update AI model versions for admin
This commit is contained in:
commit
7d2a9d1d48
14 changed files with 479 additions and 95 deletions
|
|
@ -166,6 +166,15 @@ def admin_routes(config):
|
|||
renderer="rhodecode:templates/admin/settings/settings.mako",
|
||||
)
|
||||
|
||||
config.add_route(name="admin_settings_ai_update_models", pattern="/settings/ai/model/version")
|
||||
config.add_view(
|
||||
AdminAiView,
|
||||
attr="admin_settings_ai_update_models",
|
||||
route_name="admin_settings_ai_update_models",
|
||||
request_method="POST",
|
||||
renderer="json_ext",
|
||||
)
|
||||
|
||||
config.add_route(name="admin_settings_vcs_svn_generate_cfg", pattern="/settings/vcs/svn_generate_cfg")
|
||||
config.add_view(
|
||||
AdminSvnConfigView,
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import json
|
||||
import logging
|
||||
|
||||
import formencode
|
||||
|
|
@ -5,12 +6,13 @@ import formencode
|
|||
from pyramid.httpexceptions import HTTPFound
|
||||
from rhodecode.apps._base import BaseAppView
|
||||
from rhodecode.apps._base.navigation import navigation_list
|
||||
from rhodecode.apps.ai_agents.ai_settings import AIModelName, GPTVersion, ClaudeVersion, GeminiVersion
|
||||
from rhodecode.apps.ai_agents.models.base import AIServiceBase
|
||||
from rhodecode.apps.ai_agents.ai_service import get_ai_service
|
||||
from rhodecode.apps.ai_agents.ai_settings import AIModelName
|
||||
from rhodecode.apps.ai_agents.models.base import AIService
|
||||
from rhodecode.lib.auth import LoginRequired, HasPermissionAllDecorator, CSRFRequired
|
||||
from rhodecode.lib import helpers as h
|
||||
from rhodecode.model.db import User
|
||||
from rhodecode.model.forms import AiSettingsForm
|
||||
from rhodecode.model.forms import AiSettingsForm, AiSettingsModelVersionUpdateForm
|
||||
from rhodecode.model.settings import SettingsModel
|
||||
from rhodecode.model.meta import Session
|
||||
|
||||
|
|
@ -31,22 +33,55 @@ class AdminAiView(BaseAppView):
|
|||
|
||||
app_settings = c.rc_config
|
||||
c.selected_ai_model = app_settings.get("rhodecode_ai_model", AIModelName.GPT.name)
|
||||
c.selected_ai_model_version = app_settings.get("rhodecode_ai_model_version", GPTVersion.V5_nano.name)
|
||||
c.selected_ai_model_version = app_settings.get("rhodecode_ai_model_version")
|
||||
c.api_key = app_settings.get("rhodecode_ai_api_key")
|
||||
c.ai_features_enabled = app_settings.get("rhodecode_ai_features_enabled", False)
|
||||
|
||||
c.ai_instructions = app_settings.get("rhodecode_ai_code_review_instructions")
|
||||
if not c.ai_instructions:
|
||||
# to not run formatting each time
|
||||
c.ai_instructions = "\n".join(AIServiceBase.DEFAULT_BASIC_REVIEW_POINTS)
|
||||
c.ai_instructions = "\n".join(AIService.DEFAULT_BASIC_REVIEW_POINTS)
|
||||
|
||||
gpt_versions = app_settings.get(f"rhodecode_ai_model_{AIModelName.GPT.name}_versions", "[]")
|
||||
claude_versions = app_settings.get(f"rhodecode_ai_model_{AIModelName.Claude.name}_versions", "[]")
|
||||
gemini_versions = app_settings.get(f"rhodecode_ai_model_{AIModelName.Gemini.name}_versions", "[]")
|
||||
|
||||
c.model_map = {
|
||||
AIModelName.GPT.name: [v.name for v in GPTVersion],
|
||||
AIModelName.Claude.name: [v.name for v in ClaudeVersion],
|
||||
AIModelName.Gemini.name: [v.name for v in GeminiVersion],
|
||||
AIModelName.GPT.name: json.loads(gpt_versions),
|
||||
AIModelName.Claude.name: json.loads(claude_versions),
|
||||
AIModelName.Gemini.name: json.loads(gemini_versions),
|
||||
}
|
||||
return self._get_template_context(c)
|
||||
|
||||
@LoginRequired()
|
||||
@HasPermissionAllDecorator("hg.admin")
|
||||
def admin_settings_ai_update_models(self):
|
||||
_ = self.request.translate
|
||||
versions = []
|
||||
try:
|
||||
data = self._parse_form(_, form_class=AiSettingsModelVersionUpdateForm)
|
||||
model = data["rhodecode_ai_model"]
|
||||
api_key = data["rhodecode_ai_api_key"]
|
||||
|
||||
ai_service = get_ai_service(api_key, model_name=model)
|
||||
versions = ai_service.list_model_api_and_display_names()
|
||||
if versions:
|
||||
message, category = self._save_model_versions(_, model=model, versions=versions)
|
||||
else:
|
||||
message = _("AI model API returned empty version list")
|
||||
category = "warning"
|
||||
except Exception as e:
|
||||
log.exception("Exception updating AI model versions: %s", e)
|
||||
category = "error"
|
||||
message = _(f"Error occurred during updating AI model versions, error: {str(e)}")
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"message": message,
|
||||
"message_category": category,
|
||||
"versions": versions,
|
||||
}
|
||||
|
||||
@CSRFRequired()
|
||||
@LoginRequired()
|
||||
@HasPermissionAllDecorator("hg.admin")
|
||||
|
|
@ -58,7 +93,20 @@ class AdminAiView(BaseAppView):
|
|||
data = self._parse_form(_)
|
||||
self._save_settings(_, data)
|
||||
|
||||
raise HTTPFound(h.route_path("admin_settings_ai"))
|
||||
return HTTPFound(h.route_path("admin_settings_ai"))
|
||||
|
||||
def _save_model_versions(self, _, model, versions):
|
||||
try:
|
||||
setting = f"ai_model_{model}_versions"
|
||||
str_list = json.dumps(versions)
|
||||
sett = SettingsModel().create_or_update_setting(setting, str_list)
|
||||
Session().add(sett)
|
||||
Session().commit()
|
||||
SettingsModel().invalidate_settings_cache()
|
||||
return _("AI model versions updated"), "success"
|
||||
except Exception as e:
|
||||
log.exception("Exception saving AI model versions: %s", e)
|
||||
return _(f"Error occurred during saving AI model versions, error: {str(e)}"), "error"
|
||||
|
||||
def _save_settings(self, _, data):
|
||||
try:
|
||||
|
|
@ -81,7 +129,7 @@ class AdminAiView(BaseAppView):
|
|||
h.flash(_("AI settings saved"), category="success")
|
||||
except Exception as e:
|
||||
log.exception("Exception saving AI settings: %s", e)
|
||||
h.flash(_("Error occurred during saving AI settings"), category="error")
|
||||
h.flash(_(f"Error occurred during saving AI settings, error: {str(e)}"), category="error")
|
||||
|
||||
def _activate_deactivate_ai_user(self, data, form_key):
|
||||
log.debug("%s AI user" % "Activating" if data[form_key] else "Deactivating")
|
||||
|
|
@ -89,13 +137,12 @@ class AdminAiView(BaseAppView):
|
|||
ai_user.active = data[form_key]
|
||||
Session().add(ai_user)
|
||||
|
||||
def _parse_form(self, _):
|
||||
def _parse_form(self, _, form_class=AiSettingsForm):
|
||||
try:
|
||||
form = AiSettingsForm()()
|
||||
form = form_class()()
|
||||
data = form.to_python(self.request.POST)
|
||||
except formencode.Invalid as errors:
|
||||
log.exception("Failed to add new pattern")
|
||||
error = errors
|
||||
h.flash(_(f"Unknown error: {error}"), category="error")
|
||||
h.flash(_(f"Invalid form error: {error}"), category="error")
|
||||
raise HTTPFound(h.route_path("admin_settings_ai"))
|
||||
return data
|
||||
|
|
|
|||
|
|
@ -1,19 +1,23 @@
|
|||
from rhodecode.apps.ai_agents.ai_settings import AISettings, GPTVersion, AIModelName, ClaudeVersion, GeminiVersion
|
||||
from rhodecode.apps.ai_agents.models.base import wrap_ai_exceptions
|
||||
from typing import Optional
|
||||
|
||||
from rhodecode.apps.ai_agents.ai_settings import AISettings, AIModelName
|
||||
from rhodecode.apps.ai_agents.models.base import wrap_ai_exceptions, AIService
|
||||
from rhodecode.apps.ai_agents.models.claude import ClaudeService
|
||||
from rhodecode.apps.ai_agents.models.gemini import GeminiService
|
||||
from rhodecode.apps.ai_agents.models.gpt import GPTService
|
||||
|
||||
|
||||
@wrap_ai_exceptions
|
||||
def get_ai_service(api_key: str, model_name: str = AIModelName.GPT.name, version: str = GPTVersion.V5_nano.name):
|
||||
def get_ai_service(
|
||||
api_key: str, model_name: str = AIModelName.GPT.name, api_model_version_name: str = "5-nano"
|
||||
) -> Optional[AIService]:
|
||||
assert api_key, "API key is required"
|
||||
try:
|
||||
if AIModelName[model_name] is AIModelName.GPT:
|
||||
return GPTService(
|
||||
AISettings(
|
||||
model_name=AIModelName.GPT,
|
||||
model_version=GPTVersion[version],
|
||||
model_version=api_model_version_name,
|
||||
api_key=api_key,
|
||||
)
|
||||
)
|
||||
|
|
@ -21,7 +25,7 @@ def get_ai_service(api_key: str, model_name: str = AIModelName.GPT.name, version
|
|||
return ClaudeService(
|
||||
AISettings(
|
||||
model_name=AIModelName.Claude,
|
||||
model_version=ClaudeVersion[version],
|
||||
model_version=api_model_version_name,
|
||||
api_key=api_key,
|
||||
)
|
||||
)
|
||||
|
|
@ -29,9 +33,9 @@ def get_ai_service(api_key: str, model_name: str = AIModelName.GPT.name, version
|
|||
return GeminiService(
|
||||
AISettings(
|
||||
model_name=AIModelName.Gemini,
|
||||
model_version=GeminiVersion[version],
|
||||
model_version=api_model_version_name,
|
||||
api_key=api_key,
|
||||
)
|
||||
)
|
||||
except KeyError as ke:
|
||||
raise ValueError(f"Unknown model or version: {ke}")
|
||||
raise ValueError(f"Unknown model: {ke}")
|
||||
|
|
|
|||
|
|
@ -8,36 +8,8 @@ class AIModelName(enum.StrEnum):
|
|||
Gemini = enum.auto()
|
||||
|
||||
|
||||
class GeminiVersion(enum.StrEnum):
|
||||
V2_5_pro = "2.5-pro"
|
||||
V2_5_flash = "2.5-flash"
|
||||
V2_0_flash = "2.0-flash"
|
||||
V2_5_flash_light = "2.5-flash-lite"
|
||||
V2_0_flash_light = "2.0-flash-lite"
|
||||
|
||||
|
||||
class ClaudeVersion(enum.StrEnum):
|
||||
Opus_41 = "opus-4-1"
|
||||
Opus_4 = "opus-4"
|
||||
Sonnet_4 = "sonnet-4"
|
||||
Sonnet_37 = "3-7-sonnet"
|
||||
Haiku_35 = "3-5-haiku"
|
||||
Haiku_3 = "3-haiku"
|
||||
|
||||
|
||||
class GPTVersion(enum.StrEnum):
|
||||
V5 = "5"
|
||||
V5_mini = "5-mini"
|
||||
V5_nano = "5-nano"
|
||||
V4_1 = "4.1"
|
||||
V4_1_mini = "4.1-mini"
|
||||
V4_1_nano = "4.1-nano"
|
||||
V4o = "4o"
|
||||
V4o_mini = "4o-mini"
|
||||
|
||||
|
||||
@dataclass
|
||||
class AISettings:
|
||||
model_name: AIModelName
|
||||
model_version: enum.StrEnum
|
||||
model_version: str
|
||||
api_key: str
|
||||
|
|
|
|||
|
|
@ -65,7 +65,7 @@ def wrap_ai_exceptions(f):
|
|||
return wrapper
|
||||
|
||||
|
||||
class AIServiceBase:
|
||||
class AIService:
|
||||
DEFAULT_BASIC_REVIEW_POINTS = [
|
||||
"Correctness and edge cases (logic errors, boundary conditions, invalid inputs).",
|
||||
"Error handling and resilience (fail-fast where appropriate, clear propagation, retries/backoff, cleanup).",
|
||||
|
|
@ -80,10 +80,15 @@ class AIServiceBase:
|
|||
"Dependency & supply-chain hygiene (version constraints, provenance, minimal deps, portability).",
|
||||
"Portability & interoperability (standards compliance, platform differences, encoding/locale issues).",
|
||||
]
|
||||
REQUIRED_FUNCTIONS = [
|
||||
CustomFunctions.Ping,
|
||||
CustomFunctions.Code_review,
|
||||
]
|
||||
|
||||
def __init__(self, model_settings: AISettings):
|
||||
self.model_settings = model_settings
|
||||
self._validate_mandatory_settings()
|
||||
self._client = None
|
||||
|
||||
@wrap_ai_exceptions
|
||||
def _validate_mandatory_settings(self):
|
||||
|
|
@ -144,6 +149,38 @@ class AIServiceBase:
|
|||
parts.append("") # blank line between files
|
||||
return "\n".join(parts).rstrip()
|
||||
|
||||
def get_api_model_name(self):
|
||||
return self.model_settings.model_version
|
||||
|
||||
def list_model_api_and_display_names(self) -> list[tuple[str, str]]:
|
||||
models = self.list_models()
|
||||
res = []
|
||||
for m in models:
|
||||
if getattr(m, "id") and getattr(m, "display_name"):
|
||||
res.append((m.id, m.display_name))
|
||||
else:
|
||||
self.log.warning(f"Model {m} has no id or display_name")
|
||||
return self._filter_api_models(res)
|
||||
|
||||
def _filter_api_models(self, models: list[tuple[str, str]]) -> list[tuple[str, str]]:
|
||||
res = []
|
||||
filter_models = self._list_non_text_models()
|
||||
if not filter_models:
|
||||
return models
|
||||
|
||||
for api_name, display_name in models:
|
||||
if not any(p in api_name.lower() for p in filter_models):
|
||||
res.append((api_name, display_name))
|
||||
|
||||
return res
|
||||
|
||||
def _list_non_text_models(self) -> list[str]:
|
||||
return []
|
||||
|
||||
def list_models(self):
|
||||
assert self._client is not None, "Client is not initialized"
|
||||
return self._client.models.list()
|
||||
|
||||
def _get_response_input(self, request: PingRequest | CodeReviewRequest) -> list[dict[str, str]]:
|
||||
"""
|
||||
It extracts data from internal data classes and transforms it into an input value suitable for the model’s SDK.
|
||||
|
|
@ -210,13 +247,6 @@ class AIServiceBase:
|
|||
"Treat 'Plain Text' as docs/logs/config and suggest clarity/safety where relevant."
|
||||
)
|
||||
|
||||
def get_model_name(self, model_name=None):
|
||||
if model_name:
|
||||
return model_name
|
||||
name = self.model_settings.model_name.value
|
||||
version = self.model_settings.model_version.value
|
||||
return "%s-%s" % (name.strip().lower(), version.strip().lower())
|
||||
|
||||
def _get_function(self, name):
|
||||
"""
|
||||
this is a GPT/Claude-specific function - AKA instruction for GPT
|
||||
|
|
@ -310,5 +340,8 @@ class AIServiceBase:
|
|||
def _get_response(self, request: PingRequest | CodeReviewRequest) -> Response:
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def _transform(self, resp, model=None) -> Response:
|
||||
return Response(model=self.get_model_name(model_name=model), message=resp)
|
||||
if model is None:
|
||||
model = self.get_api_model_name()
|
||||
return Response(model=model, message=resp)
|
||||
|
|
|
|||
|
|
@ -5,11 +5,10 @@ from typing import Optional, Iterable
|
|||
from anthropic import Anthropic
|
||||
from anthropic.types import ToolParam
|
||||
|
||||
from rhodecode.apps.ai_agents.ai_settings import AISettings, AIModelName, ClaudeVersion
|
||||
from rhodecode.apps.ai_agents.ai_settings import AISettings
|
||||
from rhodecode.apps.ai_agents.models.base import (
|
||||
AIServiceBase,
|
||||
AIService,
|
||||
Response,
|
||||
AIServiceError,
|
||||
PingRequest,
|
||||
CodeReviewRequest,
|
||||
Review,
|
||||
|
|
@ -19,7 +18,7 @@ from rhodecode.apps.ai_agents.models.base import (
|
|||
from rhodecode.lib.codeblocks import DiffSet
|
||||
|
||||
|
||||
class ClaudeService(AIServiceBase):
|
||||
class ClaudeService(AIService):
|
||||
def __init__(self, model_settings: AISettings):
|
||||
super().__init__(model_settings)
|
||||
self._client = Anthropic(api_key=self.model_settings.api_key)
|
||||
|
|
@ -32,12 +31,10 @@ class ClaudeService(AIServiceBase):
|
|||
|
||||
_input = self._get_response_input(request)
|
||||
|
||||
model_full_name = self._get_current_model_version()
|
||||
|
||||
function = self._adapt_function(self._get_function(request.function_name))
|
||||
|
||||
return self._client.messages.create(
|
||||
model=model_full_name,
|
||||
model=self.get_api_model_name(),
|
||||
messages=_input,
|
||||
max_tokens=max_tokens,
|
||||
tools=[
|
||||
|
|
@ -59,17 +56,6 @@ class ClaudeService(AIServiceBase):
|
|||
return Response(model=resp.model, message=r.input)
|
||||
raise ValueError("Response has incorrect return type, can't parses it.")
|
||||
|
||||
def _get_current_model_version(self):
|
||||
models = self._client.models.list()
|
||||
semi_key = "%s-%s" % (
|
||||
self.model_settings.model_name.lower().strip(),
|
||||
self.model_settings.model_version.lower().strip(),
|
||||
)
|
||||
for model in models:
|
||||
if semi_key in model.id:
|
||||
return model.id
|
||||
raise AIServiceError("Unknown model version: %s" % semi_key)
|
||||
|
||||
def _calculate_max_tokens(self, request: PingRequest | CodeReviewRequest):
|
||||
default = 256
|
||||
|
||||
|
|
|
|||
|
|
@ -11,11 +11,11 @@ from rhodecode.apps.ai_agents.models.base import (
|
|||
CodeReviewRequest,
|
||||
CustomFunctions,
|
||||
Response,
|
||||
AIServiceBase,
|
||||
AIService,
|
||||
)
|
||||
|
||||
|
||||
class GeminiService(AIServiceBase):
|
||||
class GeminiService(AIService):
|
||||
"""
|
||||
uses google compatibility option with OpenAI library: https://ai.google.dev/gemini-api/docs/openai
|
||||
"""
|
||||
|
|
@ -28,13 +28,37 @@ class GeminiService(AIServiceBase):
|
|||
)
|
||||
self.log = logging.getLogger(GeminiService.__name__)
|
||||
|
||||
def list_model_api_and_display_names(self) -> list[tuple[str, str]]:
|
||||
res = super().list_model_api_and_display_names()
|
||||
return self._remove_models_from_api_name(res)
|
||||
|
||||
def _remove_models_from_api_name(self, api_model_names: list[tuple[str, str]]) -> list[tuple[str, str]]:
|
||||
res = []
|
||||
for api_name, display_name in api_model_names:
|
||||
res.append((api_name.replace("models/", ""), display_name))
|
||||
return res
|
||||
|
||||
def _list_non_text_models(self) -> list[str]:
|
||||
return [
|
||||
"audio",
|
||||
"lyria",
|
||||
"aqa",
|
||||
"veo",
|
||||
"learnlm",
|
||||
"imagen",
|
||||
"image",
|
||||
"embedding",
|
||||
"computer-use",
|
||||
"robotics",
|
||||
]
|
||||
|
||||
def _get_response(self, request: PingRequest | CodeReviewRequest):
|
||||
_input = self._get_response_input(request)
|
||||
|
||||
function = self._adapt_function(self._get_function(request.function_name))
|
||||
|
||||
return self._client.chat.completions.create(
|
||||
model=self.get_model_name(),
|
||||
model=self.get_api_model_name(),
|
||||
messages=_input,
|
||||
tools=[function],
|
||||
tool_choice={"type": "function", "function": {"name": request.function_name}},
|
||||
|
|
|
|||
|
|
@ -1,23 +1,81 @@
|
|||
import json
|
||||
import logging
|
||||
import re
|
||||
|
||||
from openai import OpenAI
|
||||
|
||||
from rhodecode.apps.ai_agents.ai_settings import AISettings
|
||||
from rhodecode.apps.ai_agents.models.base import AIServiceBase, PingRequest, CodeReviewRequest, CustomFunctions
|
||||
from rhodecode.apps.ai_agents.models.base import AIService, PingRequest, CodeReviewRequest, CustomFunctions
|
||||
|
||||
|
||||
class GPTService(AIServiceBase):
|
||||
class GPTService(AIService):
|
||||
def __init__(self, model_settings: AISettings):
|
||||
super().__init__(model_settings)
|
||||
self._client = OpenAI(api_key=self.model_settings.api_key)
|
||||
self.log = logging.getLogger(GPTService.__name__)
|
||||
|
||||
def list_model_api_and_display_names(self) -> list[tuple[str, str]]:
|
||||
models = self.list_models()
|
||||
res = []
|
||||
for m in models:
|
||||
if getattr(m, "id"):
|
||||
model_api_name = m.id
|
||||
res.append((model_api_name, self._to_display_name(model_api_name)))
|
||||
else:
|
||||
self.log.warning(f"Model {m} has no id")
|
||||
return self._filter_api_models(res)
|
||||
|
||||
def _list_non_text_models(self) -> list[str]:
|
||||
return [
|
||||
"sora",
|
||||
"audio",
|
||||
"dall-e",
|
||||
"realtime",
|
||||
"image",
|
||||
"video",
|
||||
"codex",
|
||||
"search",
|
||||
"babbage",
|
||||
"moderation",
|
||||
"davinci",
|
||||
"transcribe",
|
||||
"tts",
|
||||
"whisper",
|
||||
"embedding",
|
||||
]
|
||||
|
||||
def _to_display_name(self, api_model_name: str) -> str:
|
||||
s = api_model_name
|
||||
sentinel = "\x00"
|
||||
|
||||
# Protect the "dall-e" brand hyphen so it stays as "dall-e"
|
||||
s = re.sub(r"\bdall-e\b", lambda m: m.group(0).replace("-", sentinel), s, flags=re.IGNORECASE)
|
||||
|
||||
# Protect dates (YYYY-MM-DD and MM-DD)
|
||||
def _protect_date(m):
|
||||
return m.group(0).replace("-", sentinel)
|
||||
|
||||
s = re.sub(r"\b\d{4}-\d{2}-\d{2}\b", _protect_date, s) # YYYY-MM-DD
|
||||
s = re.sub(r"\b\d{2}-\d{2}\b", _protect_date, s) # MM-DD
|
||||
|
||||
# Protect version-like hyphens after 2+ letters when the next token looks like a version:
|
||||
# - 1 digit optionally followed by .digits and/or letters (e.g., 4, 3.5, 4o)
|
||||
# - OR exactly 2 digits (e.g., 13) but not 3+ digits (avoid years like 2025)
|
||||
s = re.sub(
|
||||
r"([A-Za-z]{2,})-(?=(?:\d(?:\.\d+)?[A-Za-z]*|\d{2}(?!\d))(?:$|[-\s]))", lambda m: m.group(1) + sentinel, s
|
||||
)
|
||||
|
||||
# Replace all remaining hyphens with spaces
|
||||
s = s.replace("-", " ")
|
||||
|
||||
# Restore protected hyphens
|
||||
return s.replace(sentinel, "-")
|
||||
|
||||
def _get_response(self, request: PingRequest | CodeReviewRequest):
|
||||
_input = self._get_response_input(request)
|
||||
|
||||
return self._client.responses.create(
|
||||
model=self.get_model_name(),
|
||||
model=self.get_api_model_name(),
|
||||
input=_input,
|
||||
tools=[
|
||||
self._get_function(request.function_name),
|
||||
|
|
|
|||
0
rhodecode/apps/ai_agents/tests/__init__.py
Normal file
0
rhodecode/apps/ai_agents/tests/__init__.py
Normal file
|
|
@ -0,0 +1,79 @@
|
|||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from rhodecode.apps.ai_agents.models.gemini import GeminiService
|
||||
from rhodecode.apps.ai_agents.models.gpt import GPTService
|
||||
|
||||
|
||||
class TestNames:
|
||||
@pytest.mark.parametrize(
|
||||
"input_, expected_output",
|
||||
[
|
||||
("gpt-4-0613", "gpt-4 0613"),
|
||||
("gpt-4", "gpt-4"),
|
||||
("gpt-3.5-turbo", "gpt-3.5 turbo"),
|
||||
("sora-2-pro", "sora-2 pro"),
|
||||
("gpt-audio-mini-2025-10-06", "gpt audio mini 2025-10-06"),
|
||||
("gpt-realtime-mini", "gpt realtime mini"),
|
||||
("gpt-realtime-mini-2025-10-06", "gpt realtime mini 2025-10-06"),
|
||||
("sora-2", "sora-2"),
|
||||
("davinci-002", "davinci 002"),
|
||||
("babbage-002", "babbage 002"),
|
||||
("gpt-3.5-turbo-instruct", "gpt-3.5 turbo instruct"),
|
||||
("gpt-3.5-turbo-instruct-0914", "gpt-3.5 turbo instruct 0914"),
|
||||
("dall-e-3", "dall-e 3"),
|
||||
("gpt-3.5-turbo-1106", "gpt-3.5 turbo 1106"),
|
||||
("gpt-4o-audio-preview", "gpt-4o audio preview"),
|
||||
("gpt-4o-realtime-preview", "gpt-4o realtime preview"),
|
||||
("omni-moderation-latest", "omni moderation latest"),
|
||||
("omni-moderation-2024-09-26", "omni moderation 2024-09-26"),
|
||||
("gpt-4o-realtime-preview-2024-12-17", "gpt-4o realtime preview 2024-12-17"),
|
||||
("gpt-4o-audio-preview-2024-12-17", "gpt-4o audio preview 2024-12-17"),
|
||||
("gpt-4o-mini-realtime-preview-2024-12-17", "gpt-4o mini realtime preview 2024-12-17"),
|
||||
("gpt-4o-mini-audio-preview-2024-12-17", "gpt-4o mini audio preview 2024-12-17"),
|
||||
("o1-2024-12-17", "o1 2024-12-17"),
|
||||
("o1", "o1"),
|
||||
("gpt-4o-mini-realtime-preview", "gpt-4o mini realtime preview"),
|
||||
("gpt-4o-mini-audio-preview", "gpt-4o mini audio preview"),
|
||||
("o3-mini", "o3 mini"),
|
||||
("o3-mini-2025-01-31", "o3 mini 2025-01-31"),
|
||||
("gpt-4o-2024-11-20", "gpt-4o 2024-11-20"),
|
||||
("gpt-4o-search-preview-2025-03-11", "gpt-4o search preview 2025-03-11"),
|
||||
("gpt-4o-search-preview", "gpt-4o search preview"),
|
||||
("gpt-4o-mini-search-preview-2025-03-11", "gpt-4o mini search preview 2025-03-11"),
|
||||
("gpt-4o-mini-search-preview", "gpt-4o mini search preview"),
|
||||
("gpt-4o-transcribe", "gpt-4o transcribe"),
|
||||
("gpt-4o-mini-transcribe", "gpt-4o mini transcribe"),
|
||||
("o1-pro-2025-03-19", "o1 pro 2025-03-19"),
|
||||
("o1-pro", "o1 pro"),
|
||||
("gpt-4o-mini-tts", "gpt-4o mini tts"),
|
||||
("o3-2025-04-16", "o3 2025-04-16"),
|
||||
("gpt-5-pro-2025-10-06", "gpt-5 pro 2025-10-06"),
|
||||
("gpt-5-pro", "gpt-5 pro"),
|
||||
("gpt-audio-mini", "gpt audio mini"),
|
||||
("gpt-3.5-turbo-16k", "gpt-3.5 turbo 16k"),
|
||||
("tts-1", "tts-1"),
|
||||
("whisper-1", "whisper-1"),
|
||||
("text-embedding-ada-002", "text embedding ada 002"),
|
||||
],
|
||||
)
|
||||
def test_gpt_display_names(self, input_, expected_output):
|
||||
service = GPTService(MagicMock())
|
||||
display_names = service._to_display_name(input_)
|
||||
assert display_names == expected_output
|
||||
|
||||
def test_gemini_api_names(self):
|
||||
service = GeminiService(MagicMock())
|
||||
input_ = [
|
||||
("models/embedding-gecko-001", "embedding gecko 001"),
|
||||
("models/gemini-2.5-pro-preview-03-25", "gemini-2.5 pro preview 03-25"),
|
||||
]
|
||||
|
||||
cleaned_names = service._remove_models_from_api_name(input_)
|
||||
|
||||
expected_output = [
|
||||
("embedding-gecko-001", "embedding gecko 001"),
|
||||
("gemini-2.5-pro-preview-03-25", "gemini-2.5 pro preview 03-25"),
|
||||
]
|
||||
assert cleaned_names == expected_output
|
||||
98
rhodecode/apps/ai_agents/tests/test_model_versions_filter.py
Normal file
98
rhodecode/apps/ai_agents/tests/test_model_versions_filter.py
Normal file
|
|
@ -0,0 +1,98 @@
|
|||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from rhodecode.apps.ai_agents.models.claude import ClaudeService
|
||||
from rhodecode.apps.ai_agents.models.gemini import GeminiService
|
||||
from rhodecode.apps.ai_agents.models.gpt import GPTService
|
||||
|
||||
|
||||
class TestModelVersionFiltering:
|
||||
def test_filter_GPT_only_text_versions(self):
|
||||
service = GPTService(MagicMock())
|
||||
input_versions = [
|
||||
("o3", "o3"),
|
||||
("o4-mini", "o4 mini"),
|
||||
("sora-2-pro", "sora-2 pro"),
|
||||
("sora-2", "sora-2"),
|
||||
("gpt-audio-mini-2025-10-06", "gpt audio mini 2025-10-06"),
|
||||
("dall-e-3", "dall-e 3"),
|
||||
("dall-e-2", "dall-e 2"),
|
||||
("gpt-realtime", "gpt realtime"),
|
||||
("gpt-realtime-2025-08-28", "gpt realtime 2025-08-28"),
|
||||
("gpt-4o-audio-preview-2025-06-03", "gpt-4o audio preview 2025-06-03"),
|
||||
("gpt-image-1", "gpt image-1"),
|
||||
("gpt-image-1-mini", "gpt image-1 mini"),
|
||||
("gpt-5-codex", "gpt-5 codex"),
|
||||
("gpt-4o-mini-search-preview-2025-03-11", "gpt-4o mini search preview 2025-03-11"),
|
||||
("gpt-4o-mini-search-preview", "gpt-4o mini search preview"),
|
||||
("babbage-002", "babbage 002"),
|
||||
("omni-moderation-latest", "omni moderation latest"),
|
||||
("omni-moderation-2024-09-26", "omni moderation 2024-09-26"),
|
||||
("davinci-002", "davinci 002"),
|
||||
("gpt-4o-transcribe", "gpt-4o transcribe"),
|
||||
("gpt-4o-mini-transcribe", "gpt-4o mini transcribe"),
|
||||
("gpt-4o-mini-tts", "gpt-4o mini tts"),
|
||||
("whisper-1", "whisper-1"),
|
||||
("text-embedding-ada-002", "text embedding ada 002"),
|
||||
]
|
||||
|
||||
res = service._filter_api_models(input_versions)
|
||||
expected_output = [
|
||||
("o3", "o3"),
|
||||
("o4-mini", "o4 mini"),
|
||||
]
|
||||
|
||||
assert res == expected_output
|
||||
|
||||
def test_filter_Gemini_only_text_versions(self):
|
||||
service = GeminiService(MagicMock())
|
||||
input_versions = [
|
||||
("gemini-2.5-pro-preview-03-25", "Gemini 2.5 Pro Preview 03-25"),
|
||||
("veo-2.0-generate-001", "Veo 2"),
|
||||
("veo-3.0-generate-preview", "Veo 3"),
|
||||
("gemini-2.5-flash-preview-native-audio-dialog", "Gemini 2.5 Flash Preview Native Audio Dialog"),
|
||||
("lyria-realtime-exp", "Lyria Realtime Experimental"),
|
||||
("aqa", "Model that performs Attributed Question Answering."),
|
||||
("learnlm-2.0-flash-experimental", "LearnLM 2.0 Flash Experimental"),
|
||||
("imagen-3.0-generate-002", "Imagen 3.0"),
|
||||
("imagen-4.0-generate-preview-06-06", "Imagen 4 (Preview)"),
|
||||
("gemini-2.5-flash-image-preview", "Nano Banana"),
|
||||
("gemini-2.5-flash-image", "Nano Banana"),
|
||||
("gemini-embedding-exp-03-07", "Gemini Embedding Experimental 03-07"),
|
||||
("gemini-embedding-exp", "Gemini Embedding Experimental"),
|
||||
("gemini-2.5-computer-use-preview-10-2025", "Gemini 2.5 Computer Use Preview 10-2025"),
|
||||
("gemini-robotics-er-1.5-preview", "Gemini Robotics-ER 1.5 Preview"),
|
||||
]
|
||||
|
||||
res = service._filter_api_models(input_versions)
|
||||
expected_output = [
|
||||
("gemini-2.5-pro-preview-03-25", "Gemini 2.5 Pro Preview 03-25"),
|
||||
]
|
||||
|
||||
assert res == expected_output
|
||||
|
||||
def test_filter_Claude_only_text_versions(self):
|
||||
service = ClaudeService(MagicMock())
|
||||
input_versions = [
|
||||
("claude-sonnet-4-5-20250929", "Claude Sonnet 4.5"),
|
||||
("claude-opus-4-1-20250805", "Claude Opus 4.1"),
|
||||
("claude-opus-4-20250514", "Claude Opus 4"),
|
||||
("claude-sonnet-4-20250514", "Claude Sonnet 4"),
|
||||
("claude-3-7-sonnet-20250219", "Claude Sonnet 3.7"),
|
||||
("claude-3-5-haiku-20241022", "Claude Haiku 3.5"),
|
||||
("claude-3-haiku-20240307", "Claude Haiku 3"),
|
||||
]
|
||||
|
||||
res = service._filter_api_models(input_versions)
|
||||
expected_output = [
|
||||
("claude-sonnet-4-5-20250929", "Claude Sonnet 4.5"),
|
||||
("claude-opus-4-1-20250805", "Claude Opus 4.1"),
|
||||
("claude-opus-4-20250514", "Claude Opus 4"),
|
||||
("claude-sonnet-4-20250514", "Claude Sonnet 4"),
|
||||
("claude-3-7-sonnet-20250219", "Claude Sonnet 3.7"),
|
||||
("claude-3-5-haiku-20241022", "Claude Haiku 3.5"),
|
||||
("claude-3-haiku-20240307", "Claude Haiku 3"),
|
||||
]
|
||||
|
||||
assert res == expected_output
|
||||
|
|
@ -513,11 +513,6 @@ def start_ai_code_review(pull_request_id):
|
|||
log.info("AI model or model version or API key is not set, review not possible.")
|
||||
return
|
||||
|
||||
service = get_ai_service(
|
||||
api_key=ai_api_key,
|
||||
model_name=ai_model,
|
||||
version=ai_model_version,
|
||||
)
|
||||
instructions = rc_settings.get_setting_by_name("ai_code_review_instructions")
|
||||
if instructions:
|
||||
instructions = instructions.app_settings_value.split("\r\n")
|
||||
|
|
@ -535,6 +530,12 @@ def start_ai_code_review(pull_request_id):
|
|||
)
|
||||
|
||||
try:
|
||||
service = get_ai_service(
|
||||
api_key=ai_api_key,
|
||||
model_name=ai_model,
|
||||
api_model_version_name=ai_model_version,
|
||||
)
|
||||
|
||||
response = service.code_review(diffset, instructions=instructions)
|
||||
_add_comments(response, pull_request, log, ai_user)
|
||||
audit_logger.store(
|
||||
|
|
@ -548,7 +549,6 @@ def start_ai_code_review(pull_request_id):
|
|||
},
|
||||
repo=pull_request.target_repo,
|
||||
)
|
||||
|
||||
except AIServiceError as e:
|
||||
log.error("AI service error: %s", e)
|
||||
audit_logger.store(
|
||||
|
|
@ -556,7 +556,7 @@ def start_ai_code_review(pull_request_id):
|
|||
user=ai_user,
|
||||
action_data={
|
||||
audit_logger.PR_ID: pull_request_id,
|
||||
audit_logger.AI_MODEL: service.get_model_name(),
|
||||
audit_logger.AI_MODEL: service.get_api_model_name(),
|
||||
"error": True,
|
||||
"error_message": str(e),
|
||||
},
|
||||
|
|
|
|||
|
|
@ -684,3 +684,13 @@ def AiSettingsForm():
|
|||
rhodecode_ai_code_review_instructions = v.UnicodeString(strip=True)
|
||||
|
||||
return _AiSettingsForm
|
||||
|
||||
|
||||
def AiSettingsModelVersionUpdateForm():
|
||||
class _AiSettingsModelVersionUpdateForm(formencode.Schema):
|
||||
allow_extra_fields = True
|
||||
rhodecode_ai_features_enabled = v.StringBoolean(if_missing=False)
|
||||
rhodecode_ai_model = v.UnicodeString(strip=True, required=True)
|
||||
rhodecode_ai_api_key = v.UnicodeString(strip=True)
|
||||
|
||||
return _AiSettingsModelVersionUpdateForm
|
||||
|
|
|
|||
|
|
@ -15,7 +15,7 @@
|
|||
</div>
|
||||
</div>
|
||||
|
||||
<div class="field" id="model">
|
||||
<div class="field" id="model" style="display: flex;">
|
||||
<div class="label label">
|
||||
<label for="model">${_('Model')}</label>
|
||||
</div>
|
||||
|
|
@ -24,6 +24,10 @@
|
|||
<option value="${m}">${m}</option>
|
||||
% endfor
|
||||
</select>
|
||||
|
||||
<button type="button" class="btn btn-primary" id="model_version_button_update" style="display: inline-flex">
|
||||
${_('Update model versions')}
|
||||
</button>
|
||||
</div>
|
||||
|
||||
<div id="model_version_slot"></div>
|
||||
|
|
@ -76,6 +80,7 @@ import json
|
|||
const $slot = $('#model_version_slot');
|
||||
const $enable = $('#rhodecode_ai_features_enabled');
|
||||
const $form = $('#ai_features_form');
|
||||
const $updateModelsBtn = $('#model_version_button_update');
|
||||
|
||||
function renderVersionField(model) {
|
||||
const versions = MODEL_MAP[model] || [];
|
||||
|
|
@ -83,10 +88,12 @@ import json
|
|||
$slot.empty();
|
||||
|
||||
if (!versions.length) {
|
||||
setAIFieldsActive($enable.is(':checked'));
|
||||
setSaveButtonActive(false);
|
||||
return;
|
||||
}
|
||||
|
||||
setSaveButtonActive(true);
|
||||
|
||||
const html = `
|
||||
<div class="field" id="model_version">
|
||||
<div class="label label">
|
||||
|
|
@ -98,9 +105,9 @@ import json
|
|||
$slot.append(html);
|
||||
|
||||
const $ver = $('#rhodecode_ai_model_version');
|
||||
for (const v of versions) {
|
||||
const isSelected = (model === SELECTED_MODEL) && (String(v) === String(SELECTED_VERSION));
|
||||
$ver.append(new Option(String(v), String(v), isSelected, isSelected));
|
||||
for (const [key, display_name] of versions) {
|
||||
const isSelected = (model === SELECTED_MODEL) && (String(key) === String(SELECTED_VERSION));
|
||||
$ver.append(new Option(display_name, key, isSelected, isSelected));
|
||||
}
|
||||
|
||||
if (!$ver.find('option:selected').length && versions.length) {
|
||||
|
|
@ -113,6 +120,12 @@ import json
|
|||
setAIFieldsActive($enable.is(':checked'));
|
||||
}
|
||||
|
||||
function setSaveButtonActive(enable) {
|
||||
$('.buttons')
|
||||
.find('input, button')
|
||||
.prop('disabled', !enable);
|
||||
}
|
||||
|
||||
function setAIFieldsActive(enabled) {
|
||||
const $affectedFields = $('.fields .field').not('#ai-features-toggle');
|
||||
|
||||
|
|
@ -147,6 +160,57 @@ import json
|
|||
setAIFieldsActive(this.checked);
|
||||
});
|
||||
|
||||
$updateModelsBtn.click(function () {
|
||||
const $apiKey = $('#rhodecode_ai_api_key input');
|
||||
|
||||
if (!$apiKey.val()) {
|
||||
$.Topic('/notifications').publish({
|
||||
message: {
|
||||
message: "API key required to update model versions.",
|
||||
level: "warning",
|
||||
force: true
|
||||
}
|
||||
});
|
||||
return;
|
||||
}
|
||||
|
||||
const url = "${h.route_path('admin_settings_ai_update_models')}";
|
||||
$.ajax({
|
||||
type: "POST",
|
||||
url: url,
|
||||
data: {
|
||||
'rhodecode_ai_model': $model.val(),
|
||||
'rhodecode_ai_api_key': $apiKey.val(),
|
||||
'rhodecode_ai_features_enabled': $enable.is(':checked')
|
||||
},
|
||||
success: function (response) {
|
||||
$.Topic('/notifications').publish({
|
||||
message: {
|
||||
message: response["message"],
|
||||
level: response["message_category"],
|
||||
force: true
|
||||
}
|
||||
});
|
||||
|
||||
if (Array.isArray(response["versions"]) && response["versions"].length !== 0) {
|
||||
MODEL_MAP[$model.val()] = response["versions"];
|
||||
renderVersionField($model.val());
|
||||
}
|
||||
},
|
||||
error: function (data, textStatus, errorThrown) {
|
||||
console.log(data);
|
||||
$.Topic('/notifications').publish({
|
||||
message: {
|
||||
message: "Error while updating models entry.\nError code {0} ({1}). URL: {2}".format(data.status, data.statusText, url),
|
||||
level: "error",
|
||||
force: true
|
||||
}
|
||||
});
|
||||
}
|
||||
});
|
||||
|
||||
});
|
||||
|
||||
$form.submit(function () {
|
||||
let $f = $(this);
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue