diff --git a/rhodecode/apps/ai_agents/ai_service.py b/rhodecode/apps/ai_agents/ai_service.py index 21364beb..cb1c4ff9 100644 --- a/rhodecode/apps/ai_agents/ai_service.py +++ b/rhodecode/apps/ai_agents/ai_service.py @@ -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, ) diff --git a/rhodecode/apps/ai_agents/models/__init__.py b/rhodecode/apps/ai_agents/models/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/rhodecode/apps/ai_agents/models/base.py b/rhodecode/apps/ai_agents/models/base.py new file mode 100644 index 00000000..8297b2ee --- /dev/null +++ b/rhodecode/apps/ai_agents/models/base.py @@ -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) diff --git a/rhodecode/apps/ai_agents/models/gpt.py b/rhodecode/apps/ai_agents/models/gpt.py new file mode 100644 index 00000000..f94da057 --- /dev/null +++ b/rhodecode/apps/ai_agents/models/gpt.py @@ -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())