add main.py
This commit is contained in:
@@ -0,0 +1,154 @@
|
|||||||
|
"""
|
||||||
|
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
|
||||||
Reference in New Issue
Block a user