poc: changes file structure
This commit is contained in:
parent
38bebf7c74
commit
fcc3147e04
4 changed files with 119 additions and 113 deletions
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
0
rhodecode/apps/ai_agents/models/__init__.py
Normal file
0
rhodecode/apps/ai_agents/models/__init__.py
Normal file
57
rhodecode/apps/ai_agents/models/base.py
Normal file
57
rhodecode/apps/ai_agents/models/base.py
Normal 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)
|
||||
58
rhodecode/apps/ai_agents/models/gpt.py
Normal file
58
rhodecode/apps/ai_agents/models/gpt.py
Normal 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())
|
||||
Loading…
Add table
Add a link
Reference in a new issue