155 lines
6.0 KiB
Python
155 lines
6.0 KiB
Python
"""
|
||
Simple hierarchical AI agent for shopping list planning.
|
||
|
||
The script demonstrates:
|
||
* Connection to a local LLM via the OpenAI compatible API.
|
||
* A tool that internally creates a sub‑agent to estimate product prices.
|
||
* A main agent that orchestrates calls to the price tool and aggregates results.
|
||
|
||
Run with:
|
||
python -m venv .venv && source .venv/bin/activate
|
||
pip install -r requirements.txt
|
||
python main.py
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import os
|
||
import json
|
||
from typing import Dict, Any, List
|
||
|
||
# LangChain imports – the exact versions are pinned in requirements.txt
|
||
from langchain_openai import ChatOpenAI
|
||
from langchain.tools import tool
|
||
from langchain.agents import create_agent
|
||
from langchain_core.messages import HumanMessage, SystemMessage
|
||
from pydantic import SecretStr
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 1. LLM configuration – local LM Studio server
|
||
# ---------------------------------------------------------------------------
|
||
LLM_MODEL = os.getenv("LM_MODEL", "gpt-4o-mini") # default model name in LM Studio
|
||
BASE_URL = os.getenv("LM_BASE_URL", "http://localhost:1234/v1")
|
||
API_KEY = SecretStr("fake") # LM Studio does not require a real key
|
||
|
||
llm = ChatOpenAI(
|
||
model=LLM_MODEL,
|
||
base_url=BASE_URL,
|
||
api_key=API_KEY,
|
||
temperature=0.7,
|
||
)
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 2. Tool that internally creates a sub‑agent to estimate price
|
||
# ---------------------------------------------------------------------------
|
||
@tool
|
||
def get_price(product: str, city: str) -> str:
|
||
"""
|
||
Estimate the price of *product* in *city*.
|
||
|
||
The function builds a tiny sub‑agent that asks the LLM for a realistic
|
||
price table. The sub‑agent is created on every call – this keeps the
|
||
implementation simple and avoids persisting state between calls.
|
||
"""
|
||
|
||
# Sub‑agent system prompt – we keep it short to reduce token usage
|
||
sub_prompt = (
|
||
f"You are a local market price estimator for {city}. Provide a single table with columns:
|
||
| Product | Price (rub.) | Store |
|
||
The product is '{product}'. Use realistic Russian prices.")
|
||
|
||
# Create the sub‑agent – it only has one tool: none, so it just replies
|
||
sub_agent = create_agent(
|
||
llm=llm,
|
||
tools=[],
|
||
system_prompt=sub_prompt,
|
||
)
|
||
|
||
# Ask the sub‑agent for a price table
|
||
response = sub_agent.invoke({"messages": [HumanMessage(content="Generate the table.")], "configurable": {}})
|
||
# The last message contains the answer
|
||
return response["messages"][-1].content.strip()
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 3. Main agent – orchestrates calls to get_price and aggregates results
|
||
# ---------------------------------------------------------------------------
|
||
main_agent = create_agent(
|
||
llm=llm,
|
||
tools=[get_price],
|
||
system_prompt="You are a helpful assistant for planning shopping lists.",
|
||
)
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 4. Helper to format the final output nicely
|
||
# ---------------------------------------------------------------------------
|
||
|
||
def aggregate_prices(products: List[str], city: str) -> Dict[str, Any]:
|
||
"""Call get_price for each product and sum up total cost.
|
||
|
||
The function returns a dictionary with keys:
|
||
- tables: list of price tables (strings)
|
||
- total: estimated total in rubles (int or float)
|
||
"""
|
||
tables = []
|
||
total = 0.0
|
||
for prod in products:
|
||
table = get_price(prod, city)
|
||
tables.append(table)
|
||
# Extract numeric price from the table – naive regex
|
||
try:
|
||
lines = table.splitlines()
|
||
if len(lines) >= 2:
|
||
row = lines[1]
|
||
parts = [p.strip() for p in row.split('|') if p.strip()]
|
||
if len(parts) >= 2:
|
||
price_str = parts[1]
|
||
# Remove non‑digits
|
||
digits = ''.join(ch for ch in price_str if ch.isdigit())
|
||
if digits:
|
||
total += float(digits)
|
||
except Exception:
|
||
pass
|
||
return {"tables": tables, "total": total}
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 5. Main entry point – parse user input and run the agent
|
||
# ---------------------------------------------------------------------------
|
||
if __name__ == "__main__":
|
||
# Example prompt – in real usage this would come from stdin or a UI
|
||
user_prompt = (
|
||
"Помоги составить список покупок: молоко, хлеб, яблоки. Я нахожусь в Казани.")
|
||
|
||
# Run the main agent
|
||
result = main_agent.invoke({"messages": [HumanMessage(content=user_prompt)], "configurable": {}})
|
||
|
||
# Print all messages – tool calls and final answer
|
||
for msg in result["messages"]:
|
||
if hasattr(msg, "content") and msg.content:
|
||
print(msg.content)
|
||
elif hasattr(msg, "tool_calls") and msg.tool_calls:
|
||
for call in msg.tool_calls:
|
||
name = call.get("name")
|
||
args = json.dumps(call.get("args"))
|
||
print(f"{name}({args})")
|
||
|
||
# Additionally show aggregated price summary (for demonstration)
|
||
# Extract products and city from the user prompt – simple split logic
|
||
try:
|
||
parts = user_prompt.split(":", 1)[1]
|
||
prod_part, city_part = parts.split(". Я нахожусь в ")
|
||
products = [p.strip() for p in prod_part.replace("составить список покупок", "").split(",") if p.strip()]
|
||
city = city_part.rstrip("")
|
||
except Exception:
|
||
products, city = [], ""
|
||
|
||
if products and city:
|
||
agg = aggregate_prices(products, city)
|
||
print("\n--- Aggregated price tables ---")
|
||
for t in agg["tables"]:
|
||
print(t + "\n")
|
||
print(f"**Итого:** ~{int(agg['total'])} руб.")
|
||
else:
|
||
print("Не удалось извлечь список продуктов и город из запроса.")
|
||
|
||
# End of file
|