poc: gemini integrated

This commit is contained in:
ievgenii vdovenko 2025-09-10 15:33:00 +02:00
parent e984966a9d
commit d25b3965b9
7 changed files with 122 additions and 33 deletions

View file

@ -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)

View file

@ -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)

View file

@ -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):

View file

@ -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)

View file

@ -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,
}

View file

@ -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())

View file

@ -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):