poc: initial implementation of GPT model api access
This commit is contained in:
parent
fb0b40532c
commit
505f10abb3
4 changed files with 156 additions and 0 deletions
|
|
@ -305,6 +305,7 @@ whoosh==2.7.4
|
|||
zope.cachedescriptors==5.1.0
|
||||
qrcode==7.4.2
|
||||
configupdater~=3.2
|
||||
openai~=1.100.2
|
||||
|
||||
## uncomment to add the debug libraries
|
||||
#-r requirements_debug.txt
|
||||
|
|
|
|||
0
rhodecode/apps/ai_agents/__init__.py
Normal file
0
rhodecode/apps/ai_agents/__init__.py
Normal file
131
rhodecode/apps/ai_agents/ai_service.py
Normal file
131
rhodecode/apps/ai_agents/ai_service.py
Normal file
|
|
@ -0,0 +1,131 @@
|
|||
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
|
||||
|
||||
|
||||
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):
|
||||
return resp
|
||||
|
||||
|
||||
@wrap_ai_exceptions
|
||||
def get_gpt_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(
|
||||
AISettings(
|
||||
model_name=AIModelName.GPT,
|
||||
model_version=version,
|
||||
api_key=api_key,
|
||||
)
|
||||
)
|
||||
|
||||
raise ValueError("Unknown model name: %s" % model_name)
|
||||
|
||||
|
||||
@dataclass
|
||||
class Response:
|
||||
model: str
|
||||
message: str
|
||||
error: bool = False
|
||||
|
||||
|
||||
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 json.loads(item.arguments)
|
||||
|
||||
def _get_function(self, name):
|
||||
"""
|
||||
this is a GPT-specific function - AKA fine-tuning prompt
|
||||
"""
|
||||
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(
|
||||
api_key="",
|
||||
version=GPTVersion.V5_nano,
|
||||
)
|
||||
print(s.ping())
|
||||
24
rhodecode/apps/ai_agents/ai_settings.py
Normal file
24
rhodecode/apps/ai_agents/ai_settings.py
Normal file
|
|
@ -0,0 +1,24 @@
|
|||
import enum
|
||||
from dataclasses import dataclass
|
||||
|
||||
|
||||
class AIModelName(enum.StrEnum):
|
||||
GPT = enum.auto()
|
||||
|
||||
|
||||
class GPTVersion(enum.StrEnum):
|
||||
V5 = "5"
|
||||
V5_mini = "5-mini"
|
||||
V5_nano = "5-nano"
|
||||
V4_1 = "4.1"
|
||||
V4_1_mini = "4.1-mini"
|
||||
V4_1_nano = "4.1-nano"
|
||||
V4o = "4o"
|
||||
V4o_mini = "4o-mini"
|
||||
|
||||
|
||||
@dataclass
|
||||
class AISettings:
|
||||
model_name: AIModelName
|
||||
model_version: enum.StrEnum
|
||||
api_key: str
|
||||
Loading…
Add table
Add a link
Reference in a new issue