diff --git a/rhodecode/apps/admin/views/ai.py b/rhodecode/apps/admin/views/ai.py index c3570e50..8c938a51 100644 --- a/rhodecode/apps/admin/views/ai.py +++ b/rhodecode/apps/admin/views/ai.py @@ -5,7 +5,7 @@ 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 +from rhodecode.apps.ai_agents.ai_settings import AIModelName, GPTVersion, ClaudeVersion, GeminiVersion from rhodecode.apps.ai_agents.models.base import AIServiceBase from rhodecode.lib.auth import LoginRequired, HasPermissionAllDecorator, CSRFRequired from rhodecode.lib import helpers as h @@ -13,7 +13,6 @@ from rhodecode.model.db import User from rhodecode.model.forms import AiSettingsForm from rhodecode.model.settings import SettingsModel from rhodecode.model.meta import Session -from rhodecode.model.user import UserModel log = logging.getLogger(__name__) @@ -44,6 +43,7 @@ class AdminAiView(BaseAppView): 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], } return self._get_template_context(c) diff --git a/rhodecode/apps/ai_agents/ai_service.py b/rhodecode/apps/ai_agents/ai_service.py index 4e270230..4da84898 100644 --- a/rhodecode/apps/ai_agents/ai_service.py +++ b/rhodecode/apps/ai_agents/ai_service.py @@ -1,6 +1,7 @@ -from rhodecode.apps.ai_agents.ai_settings import AISettings, GPTVersion, AIModelName, ClaudeVersion +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 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 @@ -24,5 +25,13 @@ def get_ai_service(api_key: str, model_name: str = AIModelName.GPT.name, version api_key=api_key, ) ) + if AIModelName[model_name] is AIModelName.Gemini: + return GeminiService( + AISettings( + model_name=AIModelName.Gemini, + model_version=GeminiVersion[version], + api_key=api_key, + ) + ) except KeyError as ke: raise ValueError("Unknown model or version: %s" % ke) diff --git a/rhodecode/apps/ai_agents/ai_settings.py b/rhodecode/apps/ai_agents/ai_settings.py index 4eeb0c9f..e32f0947 100644 --- a/rhodecode/apps/ai_agents/ai_settings.py +++ b/rhodecode/apps/ai_agents/ai_settings.py @@ -5,6 +5,15 @@ from dataclasses import dataclass class AIModelName(enum.StrEnum): GPT = enum.auto() Claude = enum.auto() + 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): diff --git a/rhodecode/apps/ai_agents/models/base.py b/rhodecode/apps/ai_agents/models/base.py index b40f0831..7dd91eba 100644 --- a/rhodecode/apps/ai_agents/models/base.py +++ b/rhodecode/apps/ai_agents/models/base.py @@ -207,8 +207,12 @@ class AIServiceBase: "Treat 'Plain Text' as docs/logs/config and suggest clarity/safety where relevant." ) - def _get_model_name(self, model_name=None): - return model_name if model_name else self._get_model_name() + 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): """ @@ -304,4 +308,4 @@ class AIServiceBase: raise NotImplementedError def _transform(self, resp, model=None) -> Response: - return Response(model=self._get_model_name(model_name=model), message=resp) + return Response(model=self.get_model_name(model_name=model), message=resp) diff --git a/rhodecode/apps/ai_agents/models/gemini.py b/rhodecode/apps/ai_agents/models/gemini.py new file mode 100644 index 00000000..c274e61f --- /dev/null +++ b/rhodecode/apps/ai_agents/models/gemini.py @@ -0,0 +1,63 @@ +import json +import logging +from copy import deepcopy + +from openai import OpenAI + +from rhodecode.apps.ai_agents.ai_settings import AISettings + +from rhodecode.apps.ai_agents.models.base import ( + PingRequest, + CodeReviewRequest, + CustomFunctions, + Response, + AIServiceBase, +) + + +class GeminiService(AIServiceBase): + """ + uses google compatibility option with OpenAI library: https://ai.google.dev/gemini-api/docs/openai + """ + + def __init__(self, model_settings: AISettings): + super().__init__(model_settings) + self._client = OpenAI( + api_key=self.model_settings.api_key, + base_url="https://generativelanguage.googleapis.com/v1beta/openai/", + ) + self.log = logging.getLogger(GeminiService.__name__) + + 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(), + messages=_input, + tools=[function], + tool_choice={"type": "function", "function": {"name": request.function_name}}, + parallel_tool_calls=False, + ) + + def _transform(self, resp): + for item in resp.choices: + self.log.debug("response: %s", item) + if item.finish_reason == "tool_calls": + for tool in item.message.tool_calls: + if tool.function.name in [CustomFunctions.Code_review, CustomFunctions.Ping]: + return Response(model=resp.model, message=json.loads(tool.function.arguments)) + + raise ValueError("No function call found in response") + + def _adapt_function(self, function_dict: dict): + original_copy = deepcopy(function_dict) + + del original_copy["type"] + del original_copy["strict"] + + return { + "type": "function", + "function": original_copy, + } diff --git a/rhodecode/apps/ai_agents/models/gpt.py b/rhodecode/apps/ai_agents/models/gpt.py index ff2c2dcd..a1e52229 100644 --- a/rhodecode/apps/ai_agents/models/gpt.py +++ b/rhodecode/apps/ai_agents/models/gpt.py @@ -4,7 +4,7 @@ import logging from openai import OpenAI from rhodecode.apps.ai_agents.ai_settings import AISettings -from rhodecode.apps.ai_agents.models.base import AIServiceBase, PingRequest, CodeReviewRequest +from rhodecode.apps.ai_agents.models.base import AIServiceBase, PingRequest, CodeReviewRequest, CustomFunctions class GPTService(AIServiceBase): @@ -17,7 +17,7 @@ class GPTService(AIServiceBase): _input = self._get_response_input(request) return self._client.responses.create( - model=self._get_model_name(), + model=self.get_model_name(), input=_input, tools=[ self._get_function(request.function_name), @@ -28,15 +28,10 @@ class GPTService(AIServiceBase): def _transform(self, resp): for item in resp.output: - if item.type == "function_call" and item.name in ["ping_pong", "return_code_review"]: + if item.type == "function_call" and item.name in [CustomFunctions.Code_review, CustomFunctions.Ping]: return super()._transform( resp=json.loads(item.arguments), model=resp.model, ) raise ValueError("No function call found in response") - - def _get_model_name(self, model_name=None): - name = self.model_settings.model_name.value - version = self.model_settings.model_version.value - return "%s-%s" % (name.strip().lower(), version.strip().lower()) diff --git a/rhodecode/lib/celerylib/tasks.py b/rhodecode/lib/celerylib/tasks.py index 2f1e1d7d..07233de8 100644 --- a/rhodecode/lib/celerylib/tasks.py +++ b/rhodecode/lib/celerylib/tasks.py @@ -33,7 +33,7 @@ from email.utils import formatdate import rhodecode from rhodecode.apps.ai_agents.ai_service import get_ai_service -from rhodecode.apps.ai_agents.models.base import Response +from rhodecode.apps.ai_agents.models.base import Response, AIServiceError from rhodecode.lib import audit_logger, diffs, codeblocks from rhodecode.lib.celerylib import get_logger, async_task, RequestContextTask, run_task from rhodecode.lib import hooks_base @@ -533,25 +533,34 @@ def start_ai_code_review(pull_request_id): repo=pull_request.target_repo, ) - response = service.code_review(diffset, instructions=instructions) + try: + response = service.code_review(diffset, instructions=instructions) + _add_comments(response, pull_request, log, ai_user) + audit_logger.store( + "ai.code-review.finish", + user=ai_user, + action_data={ + audit_logger.PR_ID: pull_request_id, + audit_logger.AI_MODEL: response.model, + "error": response.error, + "error_message": "", + }, + repo=pull_request.target_repo, + ) - error_msg = "" - if response.error: - error_msg = response.message - - audit_logger.store( - "ai.code-review.finish", - user=ai_user, - action_data={ - audit_logger.PR_ID: pull_request_id, - audit_logger.AI_MODEL: response.model, - "error": response.error, - "error_message": error_msg, - }, - repo=pull_request.target_repo, - ) - - _add_comments(response, pull_request, log, ai_user) + except AIServiceError as e: + log.error("AI service error: %s", e) + audit_logger.store( + "ai.code-review.finish", + user=ai_user, + action_data={ + audit_logger.PR_ID: pull_request_id, + audit_logger.AI_MODEL: service.get_model_name(), + "error": True, + "error_message": str(e), + }, + repo=pull_request.target_repo, + ) def _add_comments(response: Response, pull_request: PullRequest, log: Logger | Any, ai_user: User):