poc: gemini integrated
This commit is contained in:
parent
e984966a9d
commit
d25b3965b9
7 changed files with 122 additions and 33 deletions
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
63
rhodecode/apps/ai_agents/models/gemini.py
Normal file
63
rhodecode/apps/ai_agents/models/gemini.py
Normal 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,
|
||||
}
|
||||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue