File size: 2,189 Bytes
4210028
a40b220
 
3a983bd
4210028
f3c1ab7
de88be2
4210028
de88be2
 
 
 
4210028
 
 
f3c1ab7
71d6207
de88be2
 
3a983bd
36fd23e
4210028
 
 
 
f3c1ab7
4210028
 
 
 
 
 
de88be2
 
 
 
 
 
 
 
 
4210028
 
de88be2
7cdaf78
a40b220
 
 
 
 
71d6207
 
 
f3c1ab7
71d6207
 
 
 
 
a40b220
0f2e836
71d6207
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
import os
import re
import json
from typing import List, Dict
from fastapi import FastAPI
from dotenv import load_dotenv

from langchain_core.prompts import ChatPromptTemplate
from langchain_core.prompts import PromptTemplate
from langchain_core.pydantic_v1 import BaseModel, Field
from langchain_core.output_parsers import JsonOutputParser

from langchain_groq import ChatGroq
# from api.CREDS import GROQ_API_KEY

load_dotenv()

class ChatOutput(BaseModel):
    score: int = Field(description="score of chats out of 100")
    description: str = Field(description="Short description on why the score was given, and also suggest tips for user on how could he/she make it interactive and better.")
    messages: List[str] = Field(description="Create a list of 5 potential messages for the person (self) to make the chat better and more interesting. Feel free to suggest jokes or share some fun facts to make the conversation more interactive and healthy. (just strings in a list and NOTHING else)")

app = FastAPI()

GROQ_API_KEY = os.getenv("GROQ_API_KEY")
print("--> ", GROQ_API_KEY)

llm = ChatGroq(api_key=GROQ_API_KEY, model="llama3-8b-8192")
with open("./api/prompts/system_prompt.txt", "r") as f:
    system = f.read()

human = "{text}"

parser = JsonOutputParser(pydantic_object=ChatOutput)

prompt = PromptTemplate(
    template=system,
    input_variables=['user_query'],
    partial_variables={"format_instructions": parser.get_format_instructions()},
)
chain = prompt | llm | parser

@app.post("/chat/")
async def chat(input_data: ChatOutput):
    res = chain.invoke({"user_query": input_data['text']})
    if res is None:
        return {"score": 0, "description": "No response found."}

    if isinstance(res, str):
        res = re.search(r'```(.*?)```', res, re.DOTALL)
        if not res:
            res = re.search(r'---(.*?)---', res, re.DOTALL)        
            
        if res:
            res_str = res.group(1).replace("'", '"')  
        try:
            res = json.loads(res_str)  
        except json.JSONDecodeError:
            return {"score": 0, "description": "Invalid JSON format."}
    
    print(f"Response type: {type(res)} --> {res}")
    return res