Add agent.py
This commit is contained in:
@@ -0,0 +1,136 @@
|
||||
import os
|
||||
import json
|
||||
import asyncio
|
||||
from typing import TypedDict, List, Dict, Any
|
||||
|
||||
import questionary
|
||||
from langchain_openai import ChatOpenAI
|
||||
from langgraph.graph import StateGraph, START, END
|
||||
from langgraph.checkpoint.memory import InMemorySaver
|
||||
from langgraph import interrupt, Command
|
||||
|
||||
# Define the state schema according to the plan
|
||||
class StoryState(TypedDict):
|
||||
theme: str
|
||||
scene_text: str
|
||||
options: List[str]
|
||||
chosen_option: str
|
||||
ending_text: str
|
||||
|
||||
# Configure the LLM via environment variables
|
||||
openai_api_key = os.getenv("OPENAI_API_KEY")
|
||||
openai_base_url = os.getenv("OPENAI_BASE_URL")
|
||||
llm = ChatOpenAI(
|
||||
model="gpt-4o",
|
||||
temperature=0,
|
||||
openai_api_key=openai_api_key,
|
||||
base_url=openai_base_url,
|
||||
)
|
||||
|
||||
async def call_llm(prompt: str) -> str:
|
||||
"""Invoke the LLM and return the text content, retrying once on failure."""
|
||||
for attempt in range(2):
|
||||
try:
|
||||
response = await llm.invoke(prompt)
|
||||
# ChatResult contains a list of messages
|
||||
if hasattr(response, "messages") and response.messages:
|
||||
content = response.messages[0].content
|
||||
else:
|
||||
# Fallback to direct attribute
|
||||
content = getattr(response, "content", "")
|
||||
if not content:
|
||||
raise ValueError("Empty LLM response")
|
||||
return content
|
||||
except Exception as e:
|
||||
if attempt == 1:
|
||||
raise ValueError(f"LLM call failed twice: {e}") from e
|
||||
await asyncio.sleep(1)
|
||||
raise ValueError("LLM call failed")
|
||||
|
||||
async def generate_scene(state: StoryState) -> Dict[str, Any]:
|
||||
theme = state.get("theme", "Магический лес")
|
||||
prompt = (
|
||||
f"Сгенерируй сцену на тему '{theme}'.\n"
|
||||
"Ответ должен быть в формате JSON:\n"
|
||||
"{\"scene\": \"<описание>\", \"options\": [\"<вариант1>\",\"<вариант2>\",\"<вариант3>\"]}\n"
|
||||
)
|
||||
text = await call_llm(prompt)
|
||||
try:
|
||||
data = json.loads(text)
|
||||
except Exception as e:
|
||||
raise ValueError(f"Failed to parse scene JSON: {e}")
|
||||
scene = data.get("scene", "")
|
||||
options = data.get("options", [])
|
||||
if not isinstance(options, list) or len(options) != 3:
|
||||
raise ValueError("LLM must return exactly 3 options")
|
||||
return {"scene_text": scene, "options": options}
|
||||
|
||||
async def interrupt_choice(state: StoryState) -> Dict[str, Any]:
|
||||
scene = state["scene_text"]
|
||||
options = state["options"]
|
||||
payload = {
|
||||
"type": "choice",
|
||||
"question": f"{scene} Что делаем?",
|
||||
"options": options,
|
||||
}
|
||||
# Interrupt the graph, returning a special dictionary
|
||||
return interrupt(Command(resume=payload))
|
||||
|
||||
async def generate_ending(state: StoryState) -> Dict[str, Any]:
|
||||
theme = state.get("theme", "Магический лес")
|
||||
scene = state.get("scene_text", "")
|
||||
choice = state.get("chosen_option", "")
|
||||
prompt = (
|
||||
f"На основе темы '{theme}', сцены '{scene}' и выбранного варианта '{choice}',\n"
|
||||
"сгенерируй короткую концовку истории.\n"
|
||||
"Ответ должен быть только текстом."
|
||||
)
|
||||
ending = (await call_llm(prompt)).strip()
|
||||
return {"ending_text": ending}
|
||||
|
||||
# Build the graph
|
||||
builder = StateGraph(StoryState)
|
||||
builder.add_node("generate_scene", generate_scene)
|
||||
builder.add_node("interrupt_choice", interrupt_choice)
|
||||
builder.add_node("generate_ending", generate_ending)
|
||||
|
||||
builder.set_entry_point("generate_scene")
|
||||
builder.add_edge("generate_scene", "interrupt_choice")
|
||||
builder.add_edge("interrupt_choice", "generate_ending")
|
||||
builder.add_edge("generate_ending", END)
|
||||
|
||||
# Persist state across interrupts
|
||||
builder.add_persisted_state(saver=InMemorySaver())
|
||||
|
||||
graph = builder.compile()
|
||||
|
||||
async def main() -> None:
|
||||
theme = questionary.text("Введите тему истории:").ask()
|
||||
if not theme:
|
||||
theme = "Магический лес"
|
||||
initial_state: StoryState = {
|
||||
"theme": theme,
|
||||
"scene_text": "",
|
||||
"options": [],
|
||||
"chosen_option": "",
|
||||
"ending_text": "",
|
||||
}
|
||||
# Run the graph until the first interrupt
|
||||
partial_state = await graph.ainvoke(initial_state)
|
||||
if "_interrupt" in partial_state:
|
||||
payload = partial_state["_interrupt"]
|
||||
choice = questionary.select(payload["question"], choices=payload["options"]).ask()
|
||||
partial_state["chosen_option"] = choice
|
||||
thread_id = partial_state.get("thread_id")
|
||||
config = {"configurable": {"thread_id": thread_id}} if thread_id else {}
|
||||
final_state = await graph.ainvoke(partial_state, config=config)
|
||||
else:
|
||||
final_state = partial_state
|
||||
# Print the final story
|
||||
print("\n=== Итоговая история ===\n")
|
||||
print(final_state.get("scene_text", ""))
|
||||
print(f"\nВы выбрали: {final_state.get('chosen_option', '')}\n")
|
||||
print(final_state.get("ending_text", ""))
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
Reference in New Issue
Block a user