prolific / chat_simple.py
Umar001's picture
Upload folder using huggingface_hub
5b11a78 verified
Raw
History Blame Contribute Delete
1.58 kB
import getpass
import os
from langchain.chat_models import init_chat_model
from langgraph.checkpoint.memory import MemorySaver
from langgraph.graph import START, MessagesState, StateGraph
from langchain_core.messages import AIMessage
from langchain_core.messages import HumanMessage
os.environ["GOOGLE_API_KEY"] = "AIzaSyDf2mmEYdHTVRpJ8TO5MZKEQ8jr9FhEALQ"
model = init_chat_model("gemini-2.0-flash", model_provider="google_genai")
# Define a new graph
workflow = StateGraph(state_schema=MessagesState)
# Define the function that calls the model
def call_model(state: MessagesState):
response = model.invoke(state["messages"])
return {"messages": response}
# Define the (single) node in the graph
workflow.add_edge(START, "model")
workflow.add_node("model", call_model)
# Add memory
memory = MemorySaver()
app = workflow.compile(checkpointer=memory)
# config = {"configurable": {"thread_id": "abc123"}}
# query = "Hi! I'm Bob."
# input_messages = [HumanMessage(query)]
# output = app.invoke({"messages": input_messages}, config)
# output["messages"][-1].pretty_print() # output contains all messages in state
def get_response(query, config):
input_messages = [HumanMessage(query)]
output = app.invoke({"messages": input_messages}, config)
content = output["messages"][-1].content
return content
if __name__ == "__main__":
query = input("Enter your query: ")
config = {"configurable": {"thread_id": "abc123"}}
response = get_response(query, config)
print(response)