poc: changes file structure

This commit is contained in:
ievgenii vdovenko 2025-08-22 12:54:28 +02:00
parent 38bebf7c74
commit fcc3147e04
4 changed files with 119 additions and 113 deletions

View file

@ -1,67 +1,10 @@
import json
from abc import abstractmethod
from dataclasses import dataclass
from functools import wraps
from openai import OpenAI
from rhodecode.apps.ai_agents.ai_settings import AISettings, GPTVersion, AIModelName
class AIServiceError(Exception):
pass
@dataclass
class Response:
model: str
message: str
error: bool = False
def wrap_ai_exceptions(f):
@wraps(f)
def wrapper(*args, **kwargs):
try:
return f(*args, **kwargs)
except Exception as e:
raise AIServiceError(e) from e
return wrapper
class AIServiceBase:
def __init__(self, model_settings: AISettings):
self.model_settings = model_settings
self._validate_mandatory_settings()
@wrap_ai_exceptions
def _validate_mandatory_settings(self):
api_key = self.model_settings.api_key
assert api_key is not None and api_key, "API key is required"
name = self.model_settings.model_name
assert name is not None and name, "Model name is required"
version = self.model_settings.model_version
assert version is not None and version, "Model version is required"
@wrap_ai_exceptions
def ping(self):
resp = self._get_ping_response()
return self._transform(resp)
@abstractmethod
def _get_ping_response(self):
# TODO: maybe just response
pass
def _transform(self, resp, model=None):
return Response(model=model if model else self.model_settings.model_name.value, message=resp)
from rhodecode.apps.ai_agents.models.base import wrap_ai_exceptions
from rhodecode.apps.ai_agents.models.gpt import GPTService
@wrap_ai_exceptions
def get_gpt_service(api_key: str, model_name=AIModelName.GPT, version=GPTVersion.V5_nano):
def get_ai_service(api_key: str, model_name=AIModelName.GPT, version=GPTVersion.V5_nano):
assert api_key, "API key is required"
if model_name is AIModelName.GPT:
return GPTService(
@ -75,60 +18,8 @@ def get_gpt_service(api_key: str, model_name=AIModelName.GPT, version=GPTVersion
raise ValueError("Unknown model name: %s" % model_name)
class GPTService(AIServiceBase):
def __init__(self, model_settings: AISettings):
super().__init__(model_settings)
self._client = OpenAI(api_key=self.model_settings.api_key)
def _get_ping_response(self):
return self._client.responses.create(
model=self._get_json_model_name(),
input=[{"role": "user", "content": "ping"}],
tools=[
self._get_function("ping"),
],
tool_choice={"type": "function", "name": "ping_pong"},
parallel_tool_calls=False,
)
def _transform(self, resp):
for item in resp.output:
if item.type == "function_call" and item.name == "ping_pong":
return super()._transform(
resp=json.loads(item.arguments),
model=resp.model,
)
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":
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 _:
raise ValueError("Unknown function name: %s" % name)
def _get_json_model_name(self):
name = self.model_settings.model_name.value
version = self.model_settings.model_version.value
return "%s-%s" % (name.strip().lower(), version.strip().lower())
if __name__ == "__main__":
s = get_gpt_service(
s = get_ai_service(
api_key="",
version=GPTVersion.V5_nano,
)

View file

@ -0,0 +1,57 @@
from abc import abstractmethod
from dataclasses import dataclass
from functools import wraps
from rhodecode.apps.ai_agents.ai_settings import AISettings
class AIServiceError(Exception):
pass
@dataclass
class Response:
model: str
message: str
error: bool = False
def wrap_ai_exceptions(f):
@wraps(f)
def wrapper(*args, **kwargs):
try:
return f(*args, **kwargs)
except Exception as e:
raise AIServiceError(e) from e
return wrapper
class AIServiceBase:
def __init__(self, model_settings: AISettings):
self.model_settings = model_settings
self._validate_mandatory_settings()
@wrap_ai_exceptions
def _validate_mandatory_settings(self):
api_key = self.model_settings.api_key
assert api_key is not None and api_key, "API key is required"
name = self.model_settings.model_name
assert name is not None and name, "Model name is required"
version = self.model_settings.model_version
assert version is not None and version, "Model version is required"
@wrap_ai_exceptions
def ping(self):
resp = self._get_ping_response()
return self._transform(resp)
@abstractmethod
def _get_ping_response(self):
# TODO: maybe just response
pass
def _transform(self, resp, model=None):
return Response(model=model if model else self.model_settings.model_name.value, message=resp)

View file

@ -0,0 +1,58 @@
import json
from openai import OpenAI
from rhodecode.apps.ai_agents.ai_settings import AISettings
from rhodecode.apps.ai_agents.models.base import AIServiceBase
class GPTService(AIServiceBase):
def __init__(self, model_settings: AISettings):
super().__init__(model_settings)
self._client = OpenAI(api_key=self.model_settings.api_key)
def _get_ping_response(self):
return self._client.responses.create(
model=self._get_json_model_name(),
input=[{"role": "user", "content": "ping"}],
tools=[
self._get_function("ping"),
],
tool_choice={"type": "function", "name": "ping_pong"},
parallel_tool_calls=False,
)
def _transform(self, resp):
for item in resp.output:
if item.type == "function_call" and item.name == "ping_pong":
return super()._transform(
resp=json.loads(item.arguments),
model=resp.model,
)
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":
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 _:
raise ValueError("Unknown function name: %s" % name)
def _get_json_model_name(self):
name = self.model_settings.model_name.value
version = self.model_settings.model_version.value
return "%s-%s" % (name.strip().lower(), version.strip().lower())