diff --git a/.gitignore b/.gitignore index 347ea51..ba6d7f8 100644 --- a/.gitignore +++ b/.gitignore @@ -1 +1,3 @@ -.azure \ No newline at end of file +.azure +.env +__pycache__ diff --git a/src/.env b/src/.env new file mode 100644 index 0000000..8b83c51 --- /dev/null +++ b/src/.env @@ -0,0 +1,3 @@ +OPENAI_API_VERSION= +AZURE_OPENAI_ENDPOINT= +AZURE_OPENAI_DEPLOYMENT= diff --git a/src/app.py b/src/app.py new file mode 100644 index 0000000..bad8199 --- /dev/null +++ b/src/app.py @@ -0,0 +1,65 @@ +import uuid + +from azure.identity import DefaultAzureCredential, get_bearer_token_provider +from openai import AzureOpenAI +from quart import Quart, g, request +from stateStore import StateStore +from config import Config + +token_provider = get_bearer_token_provider( + DefaultAzureCredential(), "https://cognitiveservices.azure.com/.default") + +app = Quart(__name__) + +config = Config() +state_store = StateStore() + + +@app.before_request +async def load_history(): + data = await request.get_json() + messages = data.get('messages', []) + session_state = data.get('sessionState', None) + + if session_state: + try: + history = state_store.read(session_state) + messages = history + messages + except ValueError: + pass + else: + session_state = str(uuid.uuid4()) + + g.messages = [ + {"role": "system", "content": config.system_prompt}] + messages + g.session_state = session_state + + +@app.route("/api/chat", methods=['POST']) +async def chat(): + messages = g.messages + session_state = g.session_state + + client = AzureOpenAI( + api_version=config.azure_api_version, + azure_endpoint=config.azure_endpoint, + azure_ad_token_provider=token_provider, + ) + + completion = client.chat.completions.create( + model=config.azure_deployment, + messages=messages, + max_tokens=1024 + ) + choice = completion.choices[0] + responseMessage = { + "role": choice.message.role, + "content": choice.message.content + } + + state_store.save(session_state, messages + [responseMessage]) + + return {"messages": responseMessage, "sessionState": session_state} + +if __name__ == "__main__": + app.run() diff --git a/src/config.py b/src/config.py new file mode 100644 index 0000000..0440003 --- /dev/null +++ b/src/config.py @@ -0,0 +1,28 @@ +import os +from dotenv import load_dotenv + + +class Config: + def __init__(self): + load_dotenv() + self._azure_endpoint = os.environ.get("AZURE_OPENAI_ENDPOINT") + self._azure_deployment = os.environ.get("AZURE_OPENAI_DEPLOYMENT") + self._azure_api_version = os.environ.get("OPENAI_API_VERSION") + self._system_prompt = os.environ.get( + "SYSTEM_PROMPT", "You are a helpful assistant that responds succintly to questions.") + + @property + def azure_endpoint(self): + return self._azure_endpoint + + @property + def azure_deployment(self): + return self._azure_deployment + + @property + def azure_api_version(self): + return self._azure_api_version + + @property + def system_prompt(self): + return self._system_prompt diff --git a/src/requirements.txt b/src/requirements.txt new file mode 100644 index 0000000..693de82 --- /dev/null +++ b/src/requirements.txt @@ -0,0 +1,4 @@ +quart +openai +python-dotenv +azure-identity diff --git a/src/stateStore.py b/src/stateStore.py new file mode 100644 index 0000000..d192770 --- /dev/null +++ b/src/stateStore.py @@ -0,0 +1,12 @@ +class StateStore: + def __init__(self): + self.store = {} + + def read(self, key): + state = self.store.get(key) + if state is None: + raise ValueError("Not found.") + return state + + def save(self, key, state): + self.store[key] = state