poc: implements claude ping\basic integration

This commit is contained in:
ievgenii vdovenko 2025-09-08 15:35:01 +02:00
parent 59dfb4c0a0
commit 317836aded
7 changed files with 166 additions and 66 deletions

View file

@ -306,6 +306,7 @@ zope.cachedescriptors==5.1.0
qrcode==7.4.2
configupdater~=3.2
openai~=1.100.2
anthropic~=0.66.0
## uncomment to add the debug libraries
#-r requirements_debug.txt

View file

@ -43,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],
}
return self._get_template_context(c)

View file

@ -1,5 +1,6 @@
from rhodecode.apps.ai_agents.ai_settings import AISettings, GPTVersion, AIModelName
from rhodecode.apps.ai_agents.ai_settings import AISettings, GPTVersion, AIModelName, ClaudeVersion
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.gpt import GPTService
@ -15,6 +16,14 @@ def get_ai_service(api_key: str, model_name: str = AIModelName.GPT.name, version
api_key=api_key,
)
)
if AIModelName[model_name] is AIModelName.Claude:
return ClaudeService(
AISettings(
model_name=AIModelName.Claude,
model_version=ClaudeVersion[version],
api_key=api_key,
)
)
except KeyError as ke:
raise ValueError("Unknown model or version: %s" % ke)

View file

@ -8,7 +8,12 @@ class AIModelName(enum.StrEnum):
class ClaudeVersion(enum.StrEnum):
Opus = "opus"
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):

View file

@ -118,6 +118,72 @@ class AIServiceBase:
def _get_model_name(self, model_name=None):
return model_name if model_name else self._get_model_name()
def _get_function(self, name):
"""
this is a GPT/Claude-specific function - AKA instruction for GPT
doc: https://platform.openai.com/docs/guides/function-calling
Claude uses the same tools method
"""
# TODO: unify tools
match name:
case "ping_pong":
return {
"type": "function",
"name": "ping_pong",
"description": "Always return pong",
"strict": True,
"parameters": {
"type": "object",
"properties": {"response": {"type": "string", "enum": ["pong"]}},
"required": ["response"],
"additionalProperties": False,
},
}
case "return_code_review":
return {
"type": "function",
"name": "return_code_review",
"description": (
"Return structured code review per file. Each file uses per-file 1-based "
"line numbers. Only include lines that warrant a concrete suggestion."
),
"strict": True,
"parameters": {
"type": "object",
"properties": {
"response": {
"type": "array",
"items": {
"type": "object",
"properties": {
"file_name": {"type": "string"},
"review": {
"type": "array",
"items": {
"type": "object",
"properties": {
"line_number": {"type": "integer", "minimum": 1},
"line_code": {"type": "string"},
"suggestion": {"type": "string"},
},
"required": ["line_number", "line_code", "suggestion"],
"additionalProperties": False,
},
},
},
"required": ["file_name", "review"],
"additionalProperties": False,
},
}
},
"required": ["response"],
"additionalProperties": False,
},
}
case _:
raise ValueError("Unknown function name: %s" % name)
@abstractmethod
def _get_ping_request(self) -> Request:
raise NotImplementedError
@ -130,7 +196,7 @@ class AIServiceBase:
custom_instructions: Optional[str] = None,
*args,
**kwargs,
) -> list[Request]:
) -> Request:
raise NotImplementedError
@abstractmethod

View file

@ -0,0 +1,81 @@
import logging
from dataclasses import dataclass
from typing import Optional, Iterable
from anthropic import Anthropic
from anthropic.types import MessageParam, ToolParam
from rhodecode.apps.ai_agents.ai_settings import AISettings
from rhodecode.apps.ai_agents.models.base import AIServiceBase, Request, Response, AIServiceError
from rhodecode.lib.codeblocks import DiffSet
@dataclass
class ClaudePingRequest(Request):
function_name: str
role: str
class ClaudeService(AIServiceBase):
def __init__(self, model_settings: AISettings):
super().__init__(model_settings)
self._client = Anthropic(api_key=self.model_settings.api_key)
self.log = logging.getLogger(ClaudeService.__name__)
def _get_ping_request(self) -> Request:
return ClaudePingRequest(content="ping", function_name="ping_pong", role="user")
def _get_response(self, request: ClaudePingRequest):
if isinstance(request, ClaudePingRequest):
_input: MessageParam = {
"role": request.role,
"content": request.content,
}
else:
raise ValueError("Unknown request type: %s" % type(request))
model_full_name = self._get_current_model_version()
return self._client.messages.create(
model=model_full_name,
messages=[_input],
max_tokens=1024,
tools=[
self._adapt_function(self._get_function(request.function_name)),
],
)
def _adapt_function(self, function_dict: dict) -> ToolParam:
return ToolParam(
name=function_dict["name"],
description=function_dict["description"],
input_schema=function_dict["parameters"],
)
def _transform(self, resp, model=None) -> Response:
for r in resp.content:
# TODO: push function name
if r.type == "tool_use" and r.name == "ping_pong":
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 _get_review_requests(
self,
pr_diffset: DiffSet,
basic_instructions: Optional[Iterable[str]] = None,
custom_instructions: Optional[str] = None,
*args,
**kwargs,
) -> Request:
raise NotImplementedError()

View file

@ -132,69 +132,6 @@ class GPTService(AIServiceBase):
raise ValueError("No function call found in response")
def _get_function(self, name):
"""
this is a GPT-specific function - AKA instruction for GPT
doc: https://platform.openai.com/docs/guides/function-calling
"""
match name:
case "ping_pong":
return {
"type": "function",
"name": "ping_pong",
"description": "Always return pong",
"strict": True,
"parameters": {
"type": "object",
"properties": {"response": {"type": "string", "enum": ["pong"]}},
"required": ["response"],
"additionalProperties": False,
},
}
case "return_code_review":
return {
"type": "function",
"name": "return_code_review",
"description": (
"Return structured code review per file. Each file uses per-file 1-based "
"line numbers. Only include lines that warrant a concrete suggestion."
),
"strict": True,
"parameters": {
"type": "object",
"properties": {
"response": {
"type": "array",
"items": {
"type": "object",
"properties": {
"file_name": {"type": "string"},
"review": {
"type": "array",
"items": {
"type": "object",
"properties": {
"line_number": {"type": "integer", "minimum": 1},
"line_code": {"type": "string"},
"suggestion": {"type": "string"},
},
"required": ["line_number", "line_code", "suggestion"],
"additionalProperties": False,
},
},
},
"required": ["file_name", "review"],
"additionalProperties": False,
},
}
},
"required": ["response"],
"additionalProperties": False,
},
}
case _:
raise ValueError("Unknown function name: %s" % name)
def _get_model_name(self, model_name=None):
name = self.model_settings.model_name.value
version = self.model_settings.model_version.value