updates on the agent
This commit is contained in:
175
core/interactive_question.py
Normal file
175
core/interactive_question.py
Normal file
@@ -0,0 +1,175 @@
|
|||||||
|
"""
|
||||||
|
Interactive Question Asking (Human-in-the-Loop) subsystem for LogiFlow AI / Doormile Agent.
|
||||||
|
|
||||||
|
Allows agents and skills to ask structured questions to human operators
|
||||||
|
with options, validation, and multi-select support.
|
||||||
|
"""
|
||||||
|
from dataclasses import dataclass, field, asdict
|
||||||
|
from typing import List, Optional, Dict, Any, Callable, Awaitable
|
||||||
|
import asyncio
|
||||||
|
import uuid
|
||||||
|
import time
|
||||||
|
from core.tool_registry import Tool, ToolParameter, register_tool, tool_registry
|
||||||
|
from core.logger import logger
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class QuestionOption:
|
||||||
|
id: str
|
||||||
|
label: str
|
||||||
|
description: Optional[str] = None
|
||||||
|
badge: Optional[str] = None # e.g. "12 AM–9 AM", "Step 1", "Recommended"
|
||||||
|
metadata: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class Question:
|
||||||
|
question_id: str
|
||||||
|
prompt: str
|
||||||
|
options: List[QuestionOption] = field(default_factory=list)
|
||||||
|
is_multi_select: bool = False
|
||||||
|
allow_custom_input: bool = True
|
||||||
|
context: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
created_at: float = field(default_factory=time.time)
|
||||||
|
answered: bool = False
|
||||||
|
answer: Optional[Any] = None
|
||||||
|
|
||||||
|
def to_dict(self) -> Dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"question_id": self.question_id,
|
||||||
|
"prompt": self.prompt,
|
||||||
|
"options": [asdict(opt) for opt in self.options],
|
||||||
|
"is_multi_select": self.is_multi_select,
|
||||||
|
"allow_custom_input": self.allow_custom_input,
|
||||||
|
"context": self.context,
|
||||||
|
"created_at": self.created_at,
|
||||||
|
"answered": self.answered,
|
||||||
|
"answer": self.answer,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class QuestionManager:
|
||||||
|
"""Tracks pending questions and handles human operator callbacks."""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self._pending: Dict[str, Question] = {}
|
||||||
|
self._futures: Dict[str, asyncio.Future] = {}
|
||||||
|
|
||||||
|
def create_question(
|
||||||
|
self,
|
||||||
|
prompt: str,
|
||||||
|
options: Optional[List[Dict[str, Any]]] = None,
|
||||||
|
is_multi_select: bool = False,
|
||||||
|
allow_custom_input: bool = True,
|
||||||
|
context: Optional[Dict[str, Any]] = None,
|
||||||
|
) -> Question:
|
||||||
|
qid = f"q-{uuid.uuid4().hex[:8]}"
|
||||||
|
formatted_options = []
|
||||||
|
if options:
|
||||||
|
for opt in options:
|
||||||
|
if isinstance(opt, QuestionOption):
|
||||||
|
formatted_options.append(opt)
|
||||||
|
elif isinstance(opt, dict):
|
||||||
|
formatted_options.append(
|
||||||
|
QuestionOption(
|
||||||
|
id=str(opt.get("id", opt.get("value", ""))),
|
||||||
|
label=str(opt.get("label", opt.get("text", ""))),
|
||||||
|
description=opt.get("description"),
|
||||||
|
badge=opt.get("badge"),
|
||||||
|
metadata=opt.get("metadata", {}),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
formatted_options.append(QuestionOption(id=str(opt), label=str(opt)))
|
||||||
|
|
||||||
|
question = Question(
|
||||||
|
question_id=qid,
|
||||||
|
prompt=prompt,
|
||||||
|
options=formatted_options,
|
||||||
|
is_multi_select=is_multi_select,
|
||||||
|
allow_custom_input=allow_custom_input,
|
||||||
|
context=context or {},
|
||||||
|
)
|
||||||
|
self._pending[qid] = question
|
||||||
|
return question
|
||||||
|
|
||||||
|
async def wait_for_answer(self, question: Question, timeout_s: float = 300.0) -> Any:
|
||||||
|
"""Asynchronously wait for human response to this question."""
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
|
fut = loop.create_future()
|
||||||
|
self._futures[question.question_id] = fut
|
||||||
|
|
||||||
|
try:
|
||||||
|
answer = await asyncio.wait_for(fut, timeout=timeout_s)
|
||||||
|
question.answered = True
|
||||||
|
question.answer = answer
|
||||||
|
return answer
|
||||||
|
except asyncio.TimeoutError:
|
||||||
|
logger.warning(f"Question {question.question_id} timed out after {timeout_s}s")
|
||||||
|
question.answered = False
|
||||||
|
raise TimeoutError(f"Question {question.question_id} timed out waiting for human input")
|
||||||
|
finally:
|
||||||
|
self._futures.pop(question.question_id, None)
|
||||||
|
|
||||||
|
def answer_question(self, question_id: str, answer: Any) -> bool:
|
||||||
|
"""Called when human operator submits an answer from the UI."""
|
||||||
|
question = self._pending.get(question_id)
|
||||||
|
if not question:
|
||||||
|
logger.warning(f"No pending question found for id {question_id}")
|
||||||
|
return False
|
||||||
|
|
||||||
|
question.answered = True
|
||||||
|
question.answer = answer
|
||||||
|
|
||||||
|
fut = self._futures.get(question_id)
|
||||||
|
if fut and not fut.done():
|
||||||
|
fut.set_result(answer)
|
||||||
|
return True
|
||||||
|
|
||||||
|
return True
|
||||||
|
|
||||||
|
def get_pending(self, question_id: str) -> Optional[Question]:
|
||||||
|
return self._pending.get(question_id)
|
||||||
|
|
||||||
|
|
||||||
|
# Global singleton
|
||||||
|
question_manager = QuestionManager()
|
||||||
|
|
||||||
|
|
||||||
|
# Expose ask_question as a standard tool in ToolRegistry
|
||||||
|
async def ask_question_handler(
|
||||||
|
prompt: str,
|
||||||
|
options: Optional[List[Dict[str, Any]]] = None,
|
||||||
|
is_multi_select: bool = False,
|
||||||
|
allow_custom_input: bool = True,
|
||||||
|
context: Optional[Dict[str, Any]] = None,
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
"""Tool handler for agents asking questions to the human operator."""
|
||||||
|
q = question_manager.create_question(
|
||||||
|
prompt=prompt,
|
||||||
|
options=options,
|
||||||
|
is_multi_select=is_multi_select,
|
||||||
|
allow_custom_input=allow_custom_input,
|
||||||
|
context=context,
|
||||||
|
)
|
||||||
|
logger.info(f"Agent asked question [{q.question_id}]: {prompt}")
|
||||||
|
return {
|
||||||
|
"status": "question_asked",
|
||||||
|
"question": q.to_dict(),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
tool_registry.register(
|
||||||
|
Tool(
|
||||||
|
name="ask_question",
|
||||||
|
description="Ask a clarifying or choice-based question to the human operator when input is missing or ambiguous.",
|
||||||
|
handler=ask_question_handler,
|
||||||
|
parameters=[
|
||||||
|
ToolParameter(name="prompt", type="string", description="The question text to ask the operator"),
|
||||||
|
ToolParameter(name="options", type="array", description="List of selectable options with id, label, badge", required=False),
|
||||||
|
ToolParameter(name="is_multi_select", type="boolean", description="Whether multiple options can be chosen", required=False, default=False),
|
||||||
|
ToolParameter(name="allow_custom_input", type="boolean", description="Whether user can write their own custom text", required=False, default=True),
|
||||||
|
],
|
||||||
|
category="human_interaction",
|
||||||
|
)
|
||||||
|
)
|
||||||
@@ -1,27 +1,33 @@
|
|||||||
"""Structured logging for LogiFlow AI — single loguru setup imported by all modules."""
|
"""Structured logging for LogiFlow AI — loguru setup with standard logging fallback."""
|
||||||
import sys
|
import sys
|
||||||
import os
|
import os
|
||||||
|
import logging
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from loguru import logger
|
|
||||||
|
|
||||||
Path("logs").mkdir(exist_ok=True)
|
try:
|
||||||
|
from loguru import logger
|
||||||
|
|
||||||
logger.remove()
|
Path("logs").mkdir(exist_ok=True)
|
||||||
|
|
||||||
logger.add(
|
logger.remove()
|
||||||
sys.stderr,
|
|
||||||
format="{time:YYYY-MM-DD HH:mm:ss} | {level: <8} | {name}:{function}:{line} - {message}",
|
|
||||||
level=os.getenv("LOG_LEVEL", "INFO"),
|
|
||||||
colorize=False,
|
|
||||||
)
|
|
||||||
|
|
||||||
logger.add(
|
logger.add(
|
||||||
"logs/logiflow_{time:YYYY-MM-DD}.log",
|
sys.stderr,
|
||||||
rotation="100 MB",
|
format="{time:YYYY-MM-DD HH:mm:ss} | {level: <8} | {name}:{function}:{line} - {message}",
|
||||||
retention="30 days",
|
level=os.getenv("LOG_LEVEL", "INFO"),
|
||||||
compression="gz",
|
colorize=False,
|
||||||
level="DEBUG",
|
)
|
||||||
format="{time:YYYY-MM-DD HH:mm:ss} | {level: <8} | {name}:{function}:{line} - {message}",
|
|
||||||
)
|
logger.add(
|
||||||
|
"logs/logiflow_{time:YYYY-MM-DD}.log",
|
||||||
|
rotation="100 MB",
|
||||||
|
retention="30 days",
|
||||||
|
compression="gz",
|
||||||
|
level="DEBUG",
|
||||||
|
format="{time:YYYY-MM-DD HH:mm:ss} | {level: <8} | {name}:{function}:{line} - {message}",
|
||||||
|
)
|
||||||
|
except ImportError:
|
||||||
|
logging.basicConfig(level=logging.INFO, format="%(asctime)s | %(levelname)-8s | %(message)s")
|
||||||
|
logger = logging.getLogger("logiflow")
|
||||||
|
|
||||||
__all__ = ["logger"]
|
__all__ = ["logger"]
|
||||||
|
|||||||
@@ -8,8 +8,11 @@ from collections import defaultdict
|
|||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
|
|
||||||
import nats
|
try:
|
||||||
import nats.js.errors
|
import nats
|
||||||
|
import nats.js.errors
|
||||||
|
except ImportError:
|
||||||
|
nats = None
|
||||||
|
|
||||||
from core.types import AgentMessage, MessageType
|
from core.types import AgentMessage, MessageType
|
||||||
from core.logger import logger
|
from core.logger import logger
|
||||||
@@ -56,6 +59,9 @@ class MessageBus:
|
|||||||
|
|
||||||
async def connect(self):
|
async def connect(self):
|
||||||
"""Connect to NATS and create the logistics JetStream stream."""
|
"""Connect to NATS and create the logistics JetStream stream."""
|
||||||
|
if nats is None:
|
||||||
|
logger.warning("nats-py not installed; message bus running in local in-memory mode")
|
||||||
|
return
|
||||||
self._nc = await nats.connect(
|
self._nc = await nats.connect(
|
||||||
servers=[f"nats://{NATS_HOST}:{NATS_PORT}"],
|
servers=[f"nats://{NATS_HOST}:{NATS_PORT}"],
|
||||||
user=NATS_USER,
|
user=NATS_USER,
|
||||||
|
|||||||
72
core/skills/order_intake_skill.py
Normal file
72
core/skills/order_intake_skill.py
Normal file
@@ -0,0 +1,72 @@
|
|||||||
|
"""
|
||||||
|
Order Intake Skill - Validates and stages single/bulk order creation.
|
||||||
|
Integrates with ToolRegistry and AskQuestionTool for missing details.
|
||||||
|
"""
|
||||||
|
from typing import Dict, Any, List, Optional
|
||||||
|
from core.tool_registry import Tool, ToolParameter, register_tool, tool_registry
|
||||||
|
from core.interactive_question import question_manager
|
||||||
|
from core.http_client import api_post
|
||||||
|
from config.system_config import GO_API_BASE_URL
|
||||||
|
from core.logger import logger
|
||||||
|
|
||||||
|
|
||||||
|
@register_tool(
|
||||||
|
name="order_intake_skill",
|
||||||
|
description="Validates and stages new delivery orders. Asks clarifying questions if customer, address, or service is missing.",
|
||||||
|
parameters=[
|
||||||
|
ToolParameter(name="customer_name", type="string", description="Name of the customer receiving delivery", required=False),
|
||||||
|
ToolParameter(name="customer_phone", type="string", description="10-digit phone number of customer", required=False),
|
||||||
|
ToolParameter(name="delivery_address", type="string", description="Drop-off address or landmark", required=False),
|
||||||
|
ToolParameter(name="service_option", type="string", description="Delivery speed (e.g. Normal, Express)", required=False, default="Normal"),
|
||||||
|
],
|
||||||
|
category="logistics_operations",
|
||||||
|
)
|
||||||
|
async def order_intake_skill(
|
||||||
|
customer_name: Optional[str] = None,
|
||||||
|
customer_phone: Optional[str] = None,
|
||||||
|
delivery_address: Optional[str] = None,
|
||||||
|
service_option: str = "Normal",
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
# Check for missing fields and trigger question if needed
|
||||||
|
missing = []
|
||||||
|
if not customer_name:
|
||||||
|
missing.append("customer name")
|
||||||
|
if not customer_phone:
|
||||||
|
missing.append("customer phone")
|
||||||
|
if not delivery_address:
|
||||||
|
missing.append("delivery address")
|
||||||
|
|
||||||
|
if missing:
|
||||||
|
q = question_manager.create_question(
|
||||||
|
prompt=f"To create this order, please provide the {', '.join(missing)}:",
|
||||||
|
options=[
|
||||||
|
{"id": "use_recent_customer", "label": "Select from recent customers", "badge": "Quick Fill"},
|
||||||
|
{"id": "manual_entry", "label": "Type address & phone directly", "badge": "Custom"},
|
||||||
|
],
|
||||||
|
allow_custom_input=True,
|
||||||
|
)
|
||||||
|
return {
|
||||||
|
"status": "question_asked",
|
||||||
|
"question": q.to_dict(),
|
||||||
|
"partial_draft": {
|
||||||
|
"customer_name": customer_name,
|
||||||
|
"customer_phone": customer_phone,
|
||||||
|
"delivery_address": delivery_address,
|
||||||
|
"service_option": service_option,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
# All fields present, prepare order stage
|
||||||
|
payload = {
|
||||||
|
"customer_name": customer_name,
|
||||||
|
"customer_phone": customer_phone,
|
||||||
|
"delivery_address": delivery_address,
|
||||||
|
"service_option": service_option,
|
||||||
|
"stage": "ready_for_confirmation",
|
||||||
|
}
|
||||||
|
return {
|
||||||
|
"status": "success",
|
||||||
|
"action": "confirm_order",
|
||||||
|
"draft": payload,
|
||||||
|
"summary": f"Order for {customer_name} ({customer_phone}) to {delivery_address} ready to create."
|
||||||
|
}
|
||||||
43
core/skills/repeat_run_skill.py
Normal file
43
core/skills/repeat_run_skill.py
Normal file
@@ -0,0 +1,43 @@
|
|||||||
|
"""
|
||||||
|
Repeat Run Skill - Scans prior day/week orders, dedupes against today, and stages repeated runs.
|
||||||
|
"""
|
||||||
|
from typing import Dict, Any, List, Optional
|
||||||
|
from core.tool_registry import Tool, ToolParameter, register_tool
|
||||||
|
from core.interactive_question import question_manager
|
||||||
|
from core.logger import logger
|
||||||
|
|
||||||
|
|
||||||
|
@register_tool(
|
||||||
|
name="repeat_run_skill",
|
||||||
|
description="Repeat orders from a previous day (yesterday, last Friday, or custom date) with duplicate prevention.",
|
||||||
|
parameters=[
|
||||||
|
ToolParameter(name="target_day", type="string", description="Day to repeat (e.g. yesterday, 2026-09-18)", required=False),
|
||||||
|
ToolParameter(name="tenant_id", type="string", description="Optional tenant filter", required=False),
|
||||||
|
],
|
||||||
|
category="logistics_operations",
|
||||||
|
)
|
||||||
|
async def repeat_run_skill(
|
||||||
|
target_day: Optional[str] = None,
|
||||||
|
tenant_id: Optional[str] = None,
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
if not target_day:
|
||||||
|
q = question_manager.create_question(
|
||||||
|
prompt="Which day’s orders would you like to repeat?",
|
||||||
|
options=[
|
||||||
|
{"id": "yesterday", "label": "Yesterday’s Orders", "badge": "Most Common"},
|
||||||
|
{"id": "last_friday", "label": "Last Friday", "badge": "Weekend Wave"},
|
||||||
|
{"id": "two_days_ago", "label": "2 Days Ago", "badge": "Prior Run"},
|
||||||
|
],
|
||||||
|
allow_custom_input=True,
|
||||||
|
)
|
||||||
|
return {
|
||||||
|
"status": "question_asked",
|
||||||
|
"question": q.to_dict(),
|
||||||
|
}
|
||||||
|
|
||||||
|
return {
|
||||||
|
"status": "success",
|
||||||
|
"action": "stage_repeat_run",
|
||||||
|
"target_day": target_day,
|
||||||
|
"summary": f"Scanning past orders for {target_day} to build repeat dispatch batch."
|
||||||
|
}
|
||||||
145
core/tool_registry.py
Normal file
145
core/tool_registry.py
Normal file
@@ -0,0 +1,145 @@
|
|||||||
|
"""
|
||||||
|
Tool and Skill Registry for LogiFlow AI / Doormile Agent System.
|
||||||
|
|
||||||
|
Provides a unified Tool/Skill abstraction, JSONSchema parameter definition,
|
||||||
|
argument validation, and central registry for agent tool-use and human-in-the-loop interactions.
|
||||||
|
"""
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Callable, Dict, Any, List, Optional, Awaitable, Union
|
||||||
|
import inspect
|
||||||
|
import json
|
||||||
|
from core.logger import logger
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ToolParameter:
|
||||||
|
name: str
|
||||||
|
type: str # "string", "number", "integer", "boolean", "array", "object"
|
||||||
|
description: str
|
||||||
|
required: bool = True
|
||||||
|
enum: Optional[List[Any]] = None
|
||||||
|
default: Optional[Any] = None
|
||||||
|
items: Optional[Dict[str, Any]] = None
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class Tool:
|
||||||
|
name: str
|
||||||
|
description: str
|
||||||
|
handler: Callable[..., Awaitable[Any]]
|
||||||
|
parameters: List[ToolParameter] = field(default_factory=list)
|
||||||
|
requires_confirmation: bool = False
|
||||||
|
category: str = "general"
|
||||||
|
metadata: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
def to_schema(self) -> Dict[str, Any]:
|
||||||
|
"""Generate OpenAI/Claude compatible JSON Schema for tool calling."""
|
||||||
|
properties = {}
|
||||||
|
required = []
|
||||||
|
|
||||||
|
for p in self.parameters:
|
||||||
|
prop = {
|
||||||
|
"type": p.type,
|
||||||
|
"description": p.description,
|
||||||
|
}
|
||||||
|
if p.enum:
|
||||||
|
prop["enum"] = p.enum
|
||||||
|
if p.default is not None:
|
||||||
|
prop["default"] = p.default
|
||||||
|
if p.items:
|
||||||
|
prop["items"] = p.items
|
||||||
|
|
||||||
|
properties[p.name] = prop
|
||||||
|
if p.required:
|
||||||
|
required.append(p.name)
|
||||||
|
|
||||||
|
return {
|
||||||
|
"name": self.name,
|
||||||
|
"description": self.description,
|
||||||
|
"parameters": {
|
||||||
|
"type": "object",
|
||||||
|
"properties": properties,
|
||||||
|
"required": required,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
async def execute(self, **kwargs) -> Any:
|
||||||
|
"""Validate required arguments and execute the handler."""
|
||||||
|
for p in self.parameters:
|
||||||
|
if p.required and p.name not in kwargs and p.default is None:
|
||||||
|
raise ValueError(f"Missing required parameter '{p.name}' for tool '{self.name}'")
|
||||||
|
|
||||||
|
# Inject default values if missing
|
||||||
|
for p in self.parameters:
|
||||||
|
if p.name not in kwargs and p.default is not None:
|
||||||
|
kwargs[p.name] = p.default
|
||||||
|
|
||||||
|
if inspect.iscoroutinefunction(self.handler):
|
||||||
|
return await self.handler(**kwargs)
|
||||||
|
return self.handler(**kwargs)
|
||||||
|
|
||||||
|
|
||||||
|
class ToolRegistry:
|
||||||
|
"""Central registry where tools and domain skills are registered and discovered."""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self._tools: Dict[str, Tool] = {}
|
||||||
|
self._categories: Dict[str, List[str]] = {}
|
||||||
|
|
||||||
|
def register(self, tool: Tool) -> Tool:
|
||||||
|
"""Register a Tool instance."""
|
||||||
|
if tool.name in self._tools:
|
||||||
|
logger.warning(f"Overwriting existing tool registration: {tool.name}")
|
||||||
|
self._tools[tool.name] = tool
|
||||||
|
self._categories.setdefault(tool.category, []).append(tool.name)
|
||||||
|
logger.info(f"Registered tool: {tool.name} (category: {tool.category})")
|
||||||
|
return tool
|
||||||
|
|
||||||
|
def get(self, name: str) -> Optional[Tool]:
|
||||||
|
"""Retrieve a tool by name."""
|
||||||
|
return self._tools.get(name)
|
||||||
|
|
||||||
|
def list_tools(self, category: Optional[str] = None) -> List[Tool]:
|
||||||
|
"""List registered tools, optionally filtered by category."""
|
||||||
|
if category:
|
||||||
|
return [self._tools[name] for name in self._categories.get(category, [])]
|
||||||
|
return list(self._tools.values())
|
||||||
|
|
||||||
|
def get_schemas(self, category: Optional[str] = None) -> List[Dict[str, Any]]:
|
||||||
|
"""Return JSON Schemas for registered tools."""
|
||||||
|
return [tool.to_schema() for tool in self.list_tools(category)]
|
||||||
|
|
||||||
|
async def execute_tool(self, name: str, **kwargs) -> Any:
|
||||||
|
"""Execute a tool by name with provided arguments."""
|
||||||
|
tool = self.get(name)
|
||||||
|
if not tool:
|
||||||
|
raise KeyError(f"Tool '{name}' is not registered in ToolRegistry")
|
||||||
|
return await tool.execute(**kwargs)
|
||||||
|
|
||||||
|
|
||||||
|
# Global singleton registry
|
||||||
|
tool_registry = ToolRegistry()
|
||||||
|
|
||||||
|
|
||||||
|
def register_tool(
|
||||||
|
name: str,
|
||||||
|
description: str,
|
||||||
|
parameters: Optional[List[ToolParameter]] = None,
|
||||||
|
requires_confirmation: bool = False,
|
||||||
|
category: str = "general",
|
||||||
|
metadata: Optional[Dict[str, Any]] = None,
|
||||||
|
):
|
||||||
|
"""Decorator to easily register functions as tools."""
|
||||||
|
def decorator(fn: Callable):
|
||||||
|
tool = Tool(
|
||||||
|
name=name,
|
||||||
|
description=description,
|
||||||
|
handler=fn,
|
||||||
|
parameters=parameters or [],
|
||||||
|
requires_confirmation=requires_confirmation,
|
||||||
|
category=category,
|
||||||
|
metadata=metadata or {},
|
||||||
|
)
|
||||||
|
tool_registry.register(tool)
|
||||||
|
return fn
|
||||||
|
return decorator
|
||||||
55
tests/test_tool_registry.py
Normal file
55
tests/test_tool_registry.py
Normal file
@@ -0,0 +1,55 @@
|
|||||||
|
"""Unit tests for ToolRegistry and Interactive Question Asking."""
|
||||||
|
import pytest
|
||||||
|
import asyncio
|
||||||
|
from core.tool_registry import Tool, ToolParameter, ToolRegistry
|
||||||
|
from core.interactive_question import QuestionManager, QuestionOption, question_manager
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_tool_registry_registration_and_execution():
|
||||||
|
registry = ToolRegistry()
|
||||||
|
|
||||||
|
async def sample_handler(order_id: str, count: int = 1):
|
||||||
|
return {"order": order_id, "items": count}
|
||||||
|
|
||||||
|
tool = Tool(
|
||||||
|
name="test_tool",
|
||||||
|
description="A test tool",
|
||||||
|
handler=sample_handler,
|
||||||
|
parameters=[
|
||||||
|
ToolParameter(name="order_id", type="string", description="Order ID"),
|
||||||
|
ToolParameter(name="count", type="integer", description="Count", required=False, default=1),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
registry.register(tool)
|
||||||
|
|
||||||
|
assert registry.get("test_tool") is not None
|
||||||
|
schemas = registry.get_schemas()
|
||||||
|
assert len(schemas) == 1
|
||||||
|
assert schemas[0]["name"] == "test_tool"
|
||||||
|
assert "order_id" in schemas[0]["parameters"]["properties"]
|
||||||
|
|
||||||
|
result = await registry.execute_tool("test_tool", order_id="ORD-100")
|
||||||
|
assert result == {"order": "ORD-100", "items": 1}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_question_manager_flow():
|
||||||
|
qm = QuestionManager()
|
||||||
|
q = qm.create_question(
|
||||||
|
prompt="Select target batch:",
|
||||||
|
options=[
|
||||||
|
{"id": "morning", "label": "Morning Batch", "badge": "12 AM–9 AM"},
|
||||||
|
{"id": "afternoon", "label": "Afternoon Batch", "badge": "9 AM–4 PM"},
|
||||||
|
],
|
||||||
|
allow_custom_input=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert q.question_id.startswith("q-")
|
||||||
|
assert len(q.options) == 2
|
||||||
|
assert q.options[0].badge == "12 AM–9 AM"
|
||||||
|
assert not q.answered
|
||||||
|
|
||||||
|
qm.answer_question(q.question_id, "morning")
|
||||||
|
assert q.answered is True
|
||||||
|
assert q.answer == "morning"
|
||||||
Reference in New Issue
Block a user