mirror of
https://github.com/pewdiepie-archdaemon/odysseus.git
synced 2026-10-08 07:52:20 +02:00
Squash Odysseus development history
This commit is contained in:
+169
-34
@@ -7,6 +7,7 @@ Holds the manage_calendar tool (CalDAV-backed event CRUD).
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Dict, Optional
|
||||
|
||||
from src.tools._common import _parse_tool_args
|
||||
@@ -18,7 +19,6 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
async def do_manage_calendar(content: str, owner: Optional[str] = None) -> Dict:
|
||||
"""Handle manage_calendar tool calls: list/create/update/delete calendar events (local SQLite)."""
|
||||
from datetime import datetime, timedelta
|
||||
from core.database import SessionLocal, CalendarCal, CalendarEvent, Note
|
||||
from routes.calendar_routes import (
|
||||
_ensure_default_calendar,
|
||||
@@ -28,6 +28,9 @@ async def do_manage_calendar(content: str, owner: Optional[str] = None) -> Dict:
|
||||
_resolve_base_uid,
|
||||
_push_caldav_event_after_commit,
|
||||
_record_caldav_delete_tombstone,
|
||||
_delete_calendar_reminders_for_event,
|
||||
_calendar_reminder_for_event,
|
||||
_event_to_dict,
|
||||
)
|
||||
import uuid as _uuid
|
||||
|
||||
@@ -99,13 +102,31 @@ async def do_manage_calendar(content: str, owner: Optional[str] = None) -> Dict:
|
||||
q = q.filter(CalendarCal.owner == owner)
|
||||
return q
|
||||
|
||||
def _first_present_arg(raw_args, *names: str):
|
||||
for name in names:
|
||||
if name in raw_args and raw_args.get(name) is not None:
|
||||
return raw_args.get(name)
|
||||
return None
|
||||
|
||||
def _has_reminder_request(raw_args) -> bool:
|
||||
if any(name in raw_args for name in (
|
||||
"reminder_minutes",
|
||||
"remind_before_minutes",
|
||||
"alarm_minutes",
|
||||
"reminder",
|
||||
"alarm",
|
||||
)):
|
||||
return True
|
||||
return bool(re.search(r"\b(remind|reminder|alarm)\b", str(raw_args.get("description") or ""), re.I))
|
||||
|
||||
def _reminder_minutes(raw_args) -> Optional[int]:
|
||||
raw = (
|
||||
raw_args.get("reminder_minutes")
|
||||
or raw_args.get("remind_before_minutes")
|
||||
or raw_args.get("alarm_minutes")
|
||||
or raw_args.get("reminder")
|
||||
or raw_args.get("alarm")
|
||||
raw = _first_present_arg(
|
||||
raw_args,
|
||||
"reminder_minutes",
|
||||
"remind_before_minutes",
|
||||
"alarm_minutes",
|
||||
"reminder",
|
||||
"alarm",
|
||||
)
|
||||
if raw in (None, ""):
|
||||
desc = str(raw_args.get("description") or "")
|
||||
@@ -145,6 +166,29 @@ async def do_manage_calendar(content: str, owner: Optional[str] = None) -> Dict:
|
||||
"""Parse agent event datetimes in the user's timezone when available."""
|
||||
return _parse_dt_pair(parse_due_for_user(raw))
|
||||
|
||||
def _parse_all_day_event_dt(raw: str) -> tuple[datetime, bool]:
|
||||
"""Preserve literal calendar dates for all-day events.
|
||||
|
||||
A date-only all-day value like ``2026-10-24`` is not an instant in UTC;
|
||||
it is the user's calendar day. Routing it through parse_due_for_user()
|
||||
shifts the stored naive datetime for positive timezones and makes
|
||||
birthdays render on the previous date.
|
||||
"""
|
||||
text = str(raw or "").strip()
|
||||
if re.fullmatch(r"\d{4}-\d{2}-\d{2}", text):
|
||||
return datetime.fromisoformat(text), False
|
||||
return _parse_event_dt(text)
|
||||
|
||||
def _looks_like_timed_dt(raw) -> bool:
|
||||
text = str(raw or "").strip()
|
||||
if not text or re.fullmatch(r"\d{4}-\d{2}-\d{2}", text):
|
||||
return False
|
||||
return bool(
|
||||
re.search(r"\d{4}-\d{2}-\d{2}[T\s]\d{1,2}:\d{2}", text)
|
||||
or re.search(r"\b\d{1,2}:\d{2}\b", text)
|
||||
or re.search(r"\b\d{1,2}\s*(?:am|pm)\b", text, re.I)
|
||||
)
|
||||
|
||||
def _first_nonempty_arg(*names: str):
|
||||
for name in names:
|
||||
value = args.get(name)
|
||||
@@ -168,16 +212,16 @@ async def do_manage_calendar(content: str, owner: Optional[str] = None) -> Dict:
|
||||
loc = f" @ {location}" if location else ""
|
||||
text = f"{summary}{loc} — {start_fmt}"
|
||||
due_date = remind_at.isoformat() + ("Z" if is_utc else "")
|
||||
expected_title = f"Reminder: {summary}"
|
||||
expected_title = f"Calendar reminder: {summary}"
|
||||
existing_q = db.query(Note).filter(
|
||||
Note.archived == False, # noqa: E712
|
||||
Note.due_date == due_date,
|
||||
)
|
||||
if owner is not None:
|
||||
existing_q = existing_q.filter(Note.owner == owner)
|
||||
target_title = re.sub(r"^\s*reminder\s*:\s*", "", expected_title.strip().lower())
|
||||
target_title = re.sub(r"^\s*(?:calendar\s+)?reminder\s*:\s*", "", expected_title.strip().lower())
|
||||
for existing in existing_q.limit(25).all():
|
||||
existing_title = re.sub(r"^\s*reminder\s*:\s*", "", (existing.title or "").strip().lower())
|
||||
existing_title = re.sub(r"^\s*(?:calendar\s+)?reminder\s*:\s*", "", (existing.title or "").strip().lower())
|
||||
if existing_title == target_title:
|
||||
return existing.id, "duplicate reminder already exists"
|
||||
note = Note(
|
||||
@@ -218,12 +262,13 @@ async def do_manage_calendar(content: str, owner: Optional[str] = None) -> Dict:
|
||||
end_raw = _first_nonempty_arg(
|
||||
"end", "end_time", "end_date", "range_end", "to", "dtend", "until"
|
||||
)
|
||||
query_raw = args.get("query") or args.get("date_range") or args.get("range")
|
||||
if query_raw and (not start_raw or not end_raw):
|
||||
query_raw = args.get("query")
|
||||
range_query_raw = args.get("date_range") or args.get("range")
|
||||
if (query_raw or range_query_raw) and (not start_raw or not end_raw):
|
||||
return {
|
||||
"error": (
|
||||
"list_events needs explicit start/end ISO datetimes; "
|
||||
f"resolve the requested range ({query_raw!r}) and call manage_calendar again."
|
||||
f"resolve the requested range ({(query_raw or range_query_raw)!r}) and call manage_calendar again."
|
||||
),
|
||||
"exit_code": 1,
|
||||
}
|
||||
@@ -253,23 +298,19 @@ async def do_manage_calendar(content: str, owner: Optional[str] = None) -> Dict:
|
||||
(CalendarCal.name == calendar_filter)
|
||||
)
|
||||
rows = q.order_by(CalendarEvent.dtstart).all()
|
||||
if query_raw:
|
||||
needle = str(query_raw).strip().lower()
|
||||
if needle:
|
||||
rows = [
|
||||
ev for ev in rows
|
||||
if needle in (ev.summary or "").lower()
|
||||
or needle in (ev.description or "").lower()
|
||||
or needle in (ev.location or "").lower()
|
||||
or needle in (ev.event_type or "").lower()
|
||||
]
|
||||
events = []
|
||||
for ev in rows:
|
||||
if ev.all_day:
|
||||
s, e = ev.dtstart.strftime("%Y-%m-%d"), ev.dtend.strftime("%Y-%m-%d")
|
||||
else:
|
||||
suffix = "Z" if getattr(ev, "is_utc", False) else ""
|
||||
s, e = ev.dtstart.isoformat() + suffix, ev.dtend.isoformat() + suffix
|
||||
events.append({
|
||||
"uid": ev.uid, "summary": ev.summary or "", "dtstart": s, "dtend": e,
|
||||
"all_day": ev.all_day, "description": ev.description or "",
|
||||
"location": ev.location or "",
|
||||
"calendar": ev.calendar.name if ev.calendar else "",
|
||||
"calendar_href": ev.calendar_id,
|
||||
"event_type": ev.event_type or "",
|
||||
"importance": ev.importance or "normal",
|
||||
"rrule": ev.rrule or "",
|
||||
})
|
||||
events.append(_event_to_dict(ev, db=db, owner=owner))
|
||||
if not events:
|
||||
response_text = f"No events between {start_dt.date().isoformat()} and {end_dt.date().isoformat()}."
|
||||
else:
|
||||
@@ -285,6 +326,11 @@ async def do_manage_calendar(content: str, owner: Optional[str] = None) -> Dict:
|
||||
line += f" !{ev['importance']}"
|
||||
if ev.get("rrule"):
|
||||
line += f" repeats({ev['rrule']})"
|
||||
if ev.get("has_reminder"):
|
||||
minutes = ev.get("reminder_minutes")
|
||||
line += f" 🔔 reminder"
|
||||
if minutes is not None:
|
||||
line += f" {minutes} min before"
|
||||
if ev.get("location"):
|
||||
line += f" @ {ev['location']}"
|
||||
if ev.get("calendar"):
|
||||
@@ -330,13 +376,21 @@ async def do_manage_calendar(content: str, owner: Optional[str] = None) -> Dict:
|
||||
|
||||
all_day = bool(args.get("all_day", False))
|
||||
try:
|
||||
dtstart, dtstart_is_utc = _parse_event_dt(dtstart_str)
|
||||
dtstart, dtstart_is_utc = (
|
||||
_parse_all_day_event_dt(dtstart_str)
|
||||
if all_day
|
||||
else _parse_event_dt(dtstart_str)
|
||||
)
|
||||
except ValueError as e:
|
||||
return {"error": f"Could not parse dtstart {dtstart_str!r}: {e}", "exit_code": 1}
|
||||
dtend_raw = args.get("dtend") or args.get("end") or args.get("end_time")
|
||||
if dtend_raw:
|
||||
try:
|
||||
dtend, dtend_is_utc = _parse_event_dt(dtend_raw)
|
||||
dtend, dtend_is_utc = (
|
||||
_parse_all_day_event_dt(dtend_raw)
|
||||
if all_day
|
||||
else _parse_event_dt(dtend_raw)
|
||||
)
|
||||
dtstart_is_utc = dtstart_is_utc or dtend_is_utc
|
||||
except ValueError as e:
|
||||
return {"error": f"Could not parse dtend {dtend_raw!r}: {e}", "exit_code": 1}
|
||||
@@ -397,11 +451,16 @@ async def do_manage_calendar(content: str, owner: Optional[str] = None) -> Dict:
|
||||
)
|
||||
return {
|
||||
"response": (
|
||||
f"Event already exists: '{summary}' on {dtstart_str}"
|
||||
f"Event already exists: [{summary}](#event-{existing.uid}) on {dtstart_str}"
|
||||
+ reminder_text
|
||||
),
|
||||
"uid": existing.uid,
|
||||
"dtstart": dtstart_str,
|
||||
"all_day": bool(existing.all_day),
|
||||
"anchor": f"[{summary}](#event-{existing.uid})",
|
||||
"has_reminder": bool(reminder_note_id),
|
||||
"reminder_note_id": reminder_note_id,
|
||||
"reminder_minutes": minutes_before if reminder_note_id else None,
|
||||
"reminder_skipped_reason": reminder_skipped_reason,
|
||||
"duplicate": True,
|
||||
"exit_code": 0,
|
||||
@@ -467,14 +526,23 @@ async def do_manage_calendar(content: str, owner: Optional[str] = None) -> Dict:
|
||||
return {
|
||||
"response": f"Created event [{summary}](#event-{uid}){tag_blurb} on {dtstart_str}{reminder_blurb}",
|
||||
"uid": uid,
|
||||
"dtstart": dtstart_str,
|
||||
"all_day": bool(all_day),
|
||||
"anchor": f"[{summary}](#event-{uid})",
|
||||
"has_reminder": bool(reminder_note_id),
|
||||
"reminder_note_id": reminder_note_id,
|
||||
"reminder_minutes": minutes_before if reminder_note_id else None,
|
||||
"reminder_skipped_reason": reminder_skipped_reason,
|
||||
"exit_code": 0,
|
||||
}
|
||||
|
||||
elif action == "update_event":
|
||||
uid = args.get("uid")
|
||||
# Compact routers sometimes call the identifier field ``id`` and
|
||||
# place the event title there. Accept both forms, but resolve a
|
||||
# title only when it is unique within the owner's calendar.
|
||||
uid = args.get("uid") or args.get("id") or args.get("title")
|
||||
if not uid and args.get("summary"):
|
||||
uid = args.get("summary")
|
||||
if not uid:
|
||||
return {"error": "uid is required", "exit_code": 1}
|
||||
try:
|
||||
@@ -482,6 +550,18 @@ async def do_manage_calendar(content: str, owner: Optional[str] = None) -> Dict:
|
||||
except ValueError as e:
|
||||
return {"error": str(e), "exit_code": 1}
|
||||
ev = _event_query().filter(CalendarEvent.uid == base_uid).first()
|
||||
if not ev:
|
||||
title_matches = _event_query().filter(
|
||||
CalendarEvent.summary == str(uid).strip()
|
||||
).all()
|
||||
if len(title_matches) == 1:
|
||||
ev = title_matches[0]
|
||||
base_uid = ev.uid
|
||||
elif len(title_matches) > 1:
|
||||
return {
|
||||
"error": "Multiple events have that exact title; uid is required",
|
||||
"exit_code": 1,
|
||||
}
|
||||
if not ev:
|
||||
return {"error": f"Event {uid} not found", "exit_code": 1}
|
||||
missing_id = reserve_upload_references(
|
||||
@@ -509,10 +589,15 @@ async def do_manage_calendar(content: str, owner: Optional[str] = None) -> Dict:
|
||||
_eff_all_day = (
|
||||
args["all_day"] if args.get("all_day") is not None else ev.all_day
|
||||
)
|
||||
if args.get("all_day") is None and bool(ev.all_day) and _looks_like_timed_dt(args["dtstart"]):
|
||||
_eff_all_day = False
|
||||
ev.all_day = False
|
||||
ev.dtstart, _su = _parse_event_dt(args["dtstart"])
|
||||
ev.is_utc = bool(_su and not _eff_all_day)
|
||||
if args.get("dtend") is not None:
|
||||
ev.dtend, _eu = _parse_event_dt(args["dtend"])
|
||||
if args.get("all_day") is None and bool(ev.all_day) and _looks_like_timed_dt(args["dtend"]):
|
||||
ev.all_day = False
|
||||
if args.get("all_day") is not None:
|
||||
ev.all_day = args["all_day"]
|
||||
# Tag/category + importance updates (any of these aliases).
|
||||
@@ -526,18 +611,67 @@ async def do_manage_calendar(content: str, owner: Optional[str] = None) -> Dict:
|
||||
ev.rrule = args.get("rrule") or ""
|
||||
elif str(args.get("repeat") or "").strip().lower() in {"none", "no", "off", "false", "single"}:
|
||||
ev.rrule = ""
|
||||
|
||||
reminder_text = ""
|
||||
reminder_note_id = None
|
||||
reminder_skipped_reason = None
|
||||
minutes_before = None
|
||||
if _has_reminder_request(args):
|
||||
_delete_calendar_reminders_for_event(db, owner, ev)
|
||||
minutes_before = _reminder_minutes(args)
|
||||
if minutes_before is None:
|
||||
reminder_text = "; reminder removed"
|
||||
else:
|
||||
reminder_note_id, reminder_skipped_reason = _create_calendar_reminder(
|
||||
ev.summary or "",
|
||||
ev.location or "",
|
||||
ev.dtstart,
|
||||
bool(ev.all_day),
|
||||
minutes_before,
|
||||
bool(ev.is_utc),
|
||||
)
|
||||
if reminder_note_id:
|
||||
reminder_text = f"; reminder set {minutes_before} min before"
|
||||
else:
|
||||
reminder_text = f"; reminder not set ({reminder_skipped_reason or 'reminder time already passed'})"
|
||||
|
||||
is_caldav = ev.calendar and ev.calendar.source == "caldav"
|
||||
if is_caldav:
|
||||
ev.caldav_sync_pending = "update"
|
||||
db.commit()
|
||||
if is_caldav:
|
||||
await _push_caldav_event_after_commit(owner, base_uid, "update")
|
||||
return {"response": f"Updated event {uid}", "exit_code": 0}
|
||||
return {
|
||||
"response": f"Updated event [{ev.summary or uid}](#event-{base_uid}){reminder_text}",
|
||||
"uid": base_uid,
|
||||
"dtstart": (
|
||||
(ev.dtstart.isoformat() + ("Z" if bool(ev.is_utc) and not bool(ev.all_day) else ""))
|
||||
if ev.dtstart else None
|
||||
),
|
||||
"all_day": bool(ev.all_day),
|
||||
"anchor": f"[{ev.summary or uid}](#event-{base_uid})",
|
||||
"has_reminder": bool(reminder_note_id) or bool(_calendar_reminder_for_event(db, owner, ev)),
|
||||
"reminder_note_id": reminder_note_id,
|
||||
"reminder_minutes": minutes_before if reminder_note_id else None,
|
||||
"reminder_skipped_reason": reminder_skipped_reason,
|
||||
"exit_code": 0,
|
||||
}
|
||||
|
||||
elif action == "delete_event":
|
||||
uid = args.get("uid")
|
||||
if not uid and args.get("summary"):
|
||||
# Exact-title deletion is safe when the title is unique and
|
||||
# avoids forcing a weak router through an unnecessary list
|
||||
# round-trip. Refuse ambiguous matches.
|
||||
matches = _event_query().filter(
|
||||
CalendarEvent.summary == str(args.get("summary")).strip()
|
||||
).all()
|
||||
if len(matches) == 1:
|
||||
uid = matches[0].uid
|
||||
elif len(matches) > 1:
|
||||
return {"error": "Multiple events have that exact title; uid is required", "exit_code": 1}
|
||||
if not uid:
|
||||
return {"error": "uid is required", "exit_code": 1}
|
||||
return {"error": "uid or exact summary is required", "exit_code": 1}
|
||||
try:
|
||||
base_uid = _resolve_base_uid(uid)
|
||||
except ValueError as e:
|
||||
@@ -548,6 +682,7 @@ async def do_manage_calendar(content: str, owner: Optional[str] = None) -> Dict:
|
||||
is_caldav = ev.calendar and ev.calendar.source == "caldav" and ev.remote_href
|
||||
if is_caldav:
|
||||
_record_caldav_delete_tombstone(db, ev, owner)
|
||||
_delete_calendar_reminders_for_event(db, owner, ev)
|
||||
db.delete(ev)
|
||||
db.commit()
|
||||
if is_caldav:
|
||||
|
||||
+59
-19
@@ -31,7 +31,7 @@ async def do_resolve_contact(content: str, owner: Optional[str] = None) -> Dict:
|
||||
try:
|
||||
import asyncio
|
||||
from routes import contacts_routes as cc
|
||||
all_contacts = await asyncio.to_thread(cc._fetch_contacts)
|
||||
all_contacts = await asyncio.to_thread(cc._fetch_contacts, False, owner)
|
||||
q = name.lower()
|
||||
for c in (all_contacts or []):
|
||||
hay_name = (c.get("name") or "").lower()
|
||||
@@ -96,14 +96,35 @@ async def do_manage_contact(content: str, owner: Optional[str] = None) -> Dict:
|
||||
# them in a thread so we don't block the event loop.
|
||||
import asyncio
|
||||
try:
|
||||
if action == "list":
|
||||
rows = await asyncio.to_thread(cc._fetch_contacts, True)
|
||||
if action in ("list", "search", "find"):
|
||||
rows = await asyncio.to_thread(cc._fetch_contacts, True, owner)
|
||||
query = str(args.get("query") or args.get("name") or args.get("email") or "").strip().lower()
|
||||
if action in ("search", "find") and query:
|
||||
rows = [
|
||||
c for c in rows
|
||||
if query in str(c.get("name") or "").lower()
|
||||
or query in " ".join(c.get("emails") or []).lower()
|
||||
or query in " ".join(c.get("phones") or []).lower()
|
||||
]
|
||||
if not rows:
|
||||
return {"output": "No contacts.", "exit_code": 0}
|
||||
lines = [f"{len(rows)} contacts:"]
|
||||
for c in rows:
|
||||
visible_rows = rows if action in ("search", "find") else rows[:20]
|
||||
if len(visible_rows) < len(rows):
|
||||
lines = [f"Showing {len(visible_rows)} of {len(rows)} contacts:"]
|
||||
else:
|
||||
lines = [f"{len(rows)} contacts:"]
|
||||
for c in visible_rows:
|
||||
em = ", ".join(c.get("emails") or [])
|
||||
lines.append(f"- {c.get('name') or '(no name)'} <{em}> [uid={c.get('uid','')}]")
|
||||
if c.get('phones'):
|
||||
lines.append(' Phone: ' + ', '.join(c['phones']))
|
||||
if c.get('address'):
|
||||
lines.append(' Address: ' + str(c['address']))
|
||||
if len(visible_rows) < len(rows):
|
||||
lines.append(
|
||||
f"- ...and {len(rows) - len(visible_rows)} more; "
|
||||
"search by name for an exact match"
|
||||
)
|
||||
return {"output": "\n".join(lines), "exit_code": 0}
|
||||
|
||||
if action == "add":
|
||||
@@ -121,39 +142,58 @@ async def do_manage_contact(content: str, owner: Optional[str] = None) -> Dict:
|
||||
if not name:
|
||||
name = email.split("@")[0] if email else (phones[0] if phones else "Contact")
|
||||
# Dedupe by email or phone (same as the /add route).
|
||||
existing = await asyncio.to_thread(cc._fetch_contacts)
|
||||
existing = await asyncio.to_thread(cc._fetch_contacts, False, owner)
|
||||
for c in existing:
|
||||
if email and email.lower() in [e.lower() for e in c.get("emails", [])]:
|
||||
return {"output": f"{email} is already a contact ({c.get('name','')}).", "exit_code": 0}
|
||||
if phones and any(p in (c.get("phones") or []) for p in phones):
|
||||
return {"output": f"{phones[0]} is already a contact ({c.get('name','')}).", "exit_code": 0}
|
||||
ok = await asyncio.to_thread(cc._create_contact, name, email, address, phones)
|
||||
ok = await asyncio.to_thread(cc._create_contact, name, email, address, phones, owner)
|
||||
detail = email or ", ".join(phones) or address
|
||||
return {"output": f"{'Added' if ok else 'Failed to add'} {name} ({detail}).", "exit_code": 0 if ok else 1}
|
||||
|
||||
if action in ("update", "edit"):
|
||||
uid = (args.get("uid") or "").strip()
|
||||
name = (args.get("name") or "").strip()
|
||||
existing = await asyncio.to_thread(cc._fetch_contacts, True, owner)
|
||||
if not uid and name:
|
||||
matches = [c for c in existing if str(c.get("name") or "").strip().lower() == name.lower()]
|
||||
if len(matches) == 1:
|
||||
uid = str(matches[0].get("uid") or "")
|
||||
if not uid:
|
||||
return {"error": "uid is required for update (use action=list to find it)", "exit_code": 1}
|
||||
name = (args.get("name") or "").strip()
|
||||
emails = args.get("emails")
|
||||
if emails is None and args.get("email"):
|
||||
emails = [args["email"]]
|
||||
emails = [e.strip() for e in (emails or []) if e and e.strip()]
|
||||
phones = [p.strip() for p in (args.get("phones") or []) if p and p.strip()]
|
||||
address = (args.get("address") or "").strip()
|
||||
if not name and not emails and not phones and not address:
|
||||
current = next((c for c in existing if c.get('uid') == uid), None)
|
||||
if current is None:
|
||||
return {"error": "Contact not found", "exit_code": 1}
|
||||
if not {'name', 'emails', 'email', 'phones', 'address'}.intersection(args):
|
||||
return {"error": "Provide a name, emails, phones, or address to update", "exit_code": 1}
|
||||
if not name and emails:
|
||||
name = emails[0].split("@")[0]
|
||||
ok = await asyncio.to_thread(cc._update_contact, uid, name, emails, phones, address)
|
||||
# Tool updates are patches; the storage helper rewrites the whole
|
||||
# contact. Omitted fields must survive that conversion unchanged.
|
||||
name = name if 'name' in args else current.get('name', '')
|
||||
if 'emails' in args:
|
||||
emails = args['emails']
|
||||
elif 'email' in args:
|
||||
emails = [args['email']]
|
||||
else:
|
||||
emails = current.get('emails', [])
|
||||
emails = [e.strip() for e in (emails or []) if e and e.strip()]
|
||||
phones = args['phones'] if 'phones' in args else current.get('phones', [])
|
||||
phones = [p.strip() for p in (phones or []) if p and p.strip()]
|
||||
address = (args.get('address') or '').strip() if 'address' in args else current.get('address', '')
|
||||
ok = await asyncio.to_thread(cc._update_contact, uid, name, emails, phones, address, owner)
|
||||
return {"output": "Contact updated." if ok else "Update failed.", "exit_code": 0 if ok else 1}
|
||||
|
||||
if action == "delete":
|
||||
uid = (args.get("uid") or "").strip()
|
||||
name = (args.get("name") or "").strip()
|
||||
if not uid and name:
|
||||
matches = await asyncio.to_thread(cc._fetch_contacts, True, owner)
|
||||
matches = [c for c in matches if str(c.get("name") or "").strip().lower() == name.lower()]
|
||||
if len(matches) == 1:
|
||||
uid = str(matches[0].get("uid") or "")
|
||||
if not uid:
|
||||
return {"error": "uid is required for delete (use action=list to find it)", "exit_code": 1}
|
||||
ok = await asyncio.to_thread(cc._delete_contact, uid)
|
||||
ok = await asyncio.to_thread(cc._delete_contact, uid, owner)
|
||||
return {"output": "Contact deleted." if ok else "Delete failed.", "exit_code": 0 if ok else 1}
|
||||
|
||||
return {"error": f"Unknown action '{action}'. Use list, add, update, or delete.", "exit_code": 1}
|
||||
|
||||
+310
-41
@@ -10,6 +10,7 @@ them does a function-local import to avoid a top-level circular dependency,
|
||||
matching the system-domain split.
|
||||
"""
|
||||
import asyncio
|
||||
import contextlib
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
@@ -41,6 +42,132 @@ def _cookbook_is_exact_repo_id(value: Any) -> bool:
|
||||
return bool(re.fullmatch(r"[A-Za-z0-9_.-]+/[A-Za-z0-9_.-]+", str(value or "").strip()))
|
||||
|
||||
|
||||
_HF_OFFICIAL_AUTHOR_ALIASES: Dict[str, str] = {
|
||||
"qwen": "Qwen",
|
||||
"qwen2": "Qwen",
|
||||
"qwen3": "Qwen",
|
||||
"qwen4": "Qwen",
|
||||
"deepseek": "deepseek-ai",
|
||||
"deepseek-ai": "deepseek-ai",
|
||||
"llama": "meta-llama",
|
||||
"meta": "meta-llama",
|
||||
"meta-llama": "meta-llama",
|
||||
"mistral": "mistralai",
|
||||
"mixtral": "mistralai",
|
||||
"codestral": "mistralai",
|
||||
"mistralai": "mistralai",
|
||||
"gemma": "google",
|
||||
"google": "google",
|
||||
"phi": "microsoft",
|
||||
"microsoft": "microsoft",
|
||||
"nemotron": "nvidia",
|
||||
"nvidia": "nvidia",
|
||||
"gpt-oss": "openai",
|
||||
"openai": "openai",
|
||||
"kimi": "moonshotai",
|
||||
"moonshot": "moonshotai",
|
||||
"moonshotai": "moonshotai",
|
||||
"stable-diffusion": "stabilityai",
|
||||
"stability": "stabilityai",
|
||||
"stabilityai": "stabilityai",
|
||||
"falcon": "tiiuae",
|
||||
"tii": "tiiuae",
|
||||
"tiiuae": "tiiuae",
|
||||
"granite": "ibm-granite",
|
||||
"ibm": "ibm-granite",
|
||||
"allenai": "allenai",
|
||||
"olmo": "allenai",
|
||||
"huggingfacetb": "HuggingFaceTB",
|
||||
"smollm": "HuggingFaceTB",
|
||||
}
|
||||
|
||||
|
||||
_HF_SEARCH_STOPWORDS = {
|
||||
"a", "an", "and", "are", "best", "by", "can", "find", "for", "from",
|
||||
"hf", "hugging", "huggingface", "in", "is", "latest", "link", "me",
|
||||
"model", "models", "new", "newest", "official", "on", "out", "recent",
|
||||
"released", "search", "show", "the", "there", "to", "what", "with",
|
||||
}
|
||||
|
||||
|
||||
_HF_QUANT_TERMS = {
|
||||
"awq", "gguf", "gptq", "exl2", "mlx", "fp8", "fp4", "int8", "int4",
|
||||
"q8", "q6", "q5", "q4", "q3", "q2", "quant", "quantized", "quantization",
|
||||
"4bit", "8bit",
|
||||
}
|
||||
|
||||
|
||||
def _hf_official_author_for_query(query: str) -> Optional[str]:
|
||||
q = str(query or "").strip()
|
||||
if _cookbook_is_exact_repo_id(q):
|
||||
return q.split("/", 1)[0]
|
||||
lowered = q.lower()
|
||||
for alias, author in sorted(_HF_OFFICIAL_AUTHOR_ALIASES.items(), key=lambda item: len(item[0]), reverse=True):
|
||||
if re.search(rf"(?<![a-z0-9]){re.escape(alias)}(?![a-z0-9])", lowered):
|
||||
return author
|
||||
return None
|
||||
|
||||
|
||||
def _hf_query_mentions_quant(query: str) -> bool:
|
||||
lowered = str(query or "").lower()
|
||||
return any(re.search(rf"(?<![a-z0-9]){re.escape(term)}(?![a-z0-9])", lowered) for term in _HF_QUANT_TERMS)
|
||||
|
||||
|
||||
def _hf_query_terms(query: str, author: str = "") -> List[str]:
|
||||
lowered = str(query or "").lower()
|
||||
author_bits = {author.lower()}
|
||||
author_bits.update(k for k, v in _HF_OFFICIAL_AUTHOR_ALIASES.items() if v.lower() == author.lower())
|
||||
terms: List[str] = []
|
||||
for term in re.findall(r"[a-z0-9]+(?:\.[a-z0-9]+)?", lowered):
|
||||
if term in _HF_SEARCH_STOPWORDS or term in author_bits:
|
||||
continue
|
||||
if term not in terms:
|
||||
terms.append(term)
|
||||
return terms
|
||||
|
||||
|
||||
def _hf_row_text(row: Dict[str, Any]) -> str:
|
||||
parts = [
|
||||
row.get("id"),
|
||||
row.get("modelId"),
|
||||
row.get("pipeline_tag"),
|
||||
row.get("library_name"),
|
||||
" ".join(str(t) for t in (row.get("tags") or []) if t),
|
||||
]
|
||||
return " ".join(str(p or "") for p in parts).lower()
|
||||
|
||||
|
||||
def _hf_model_matches_terms(row: Dict[str, Any], terms: List[str]) -> bool:
|
||||
if not terms:
|
||||
return True
|
||||
haystack = _hf_row_text(row)
|
||||
return all(term in haystack for term in terms)
|
||||
|
||||
|
||||
def _hf_model_is_quant_variant(row: Dict[str, Any]) -> bool:
|
||||
haystack = _hf_row_text(row)
|
||||
return any(re.search(rf"(?<![a-z0-9]){re.escape(term)}(?![a-z0-9])", haystack) for term in _HF_QUANT_TERMS)
|
||||
|
||||
|
||||
def _hf_format_model_search_output(models: List[Dict[str, Any]], query: str, official_author: str = "") -> str:
|
||||
scope = f" official {official_author} model(s)" if official_author else " model(s)"
|
||||
lines = [f"Found {len(models)}{scope} for {query!r}:" if query else f"Found {len(models)}{scope}:"]
|
||||
for m in models:
|
||||
repo_id = str(m.get("id") or m.get("modelId") or "?")
|
||||
bits = []
|
||||
if m.get("pipeline_tag"):
|
||||
bits.append(str(m["pipeline_tag"]))
|
||||
if m.get("downloads") is not None:
|
||||
bits.append(f"{m['downloads']} downloads")
|
||||
if m.get("likes") is not None:
|
||||
bits.append(f"{m['likes']} likes")
|
||||
if m.get("lastModified"):
|
||||
bits.append(f"updated {m['lastModified']}")
|
||||
suffix = f" ({'; '.join(bits)})" if bits else ""
|
||||
lines.append(f"- {repo_id}{suffix}\n URL: https://huggingface.co/{repo_id}")
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def _cookbook_match_saved_preset(query: str, presets: List[Any], host: str = "") -> Optional[Dict[str, Any]]:
|
||||
"""Resolve a user-facing model label to a saved serve preset.
|
||||
|
||||
@@ -950,19 +1077,37 @@ async def _cookbook_kill_session(session_id: str, *, remote_host: str = "",
|
||||
target_label = session_id
|
||||
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=15) as client:
|
||||
resp = await client.post(f"{_INTERNAL_BASE}/api/shell/exec",
|
||||
json={"command": cmd}, headers=headers)
|
||||
if resp.status_code >= 400:
|
||||
return {
|
||||
"error": f"shell/exec returned HTTP {resp.status_code}: {resp.text[:200]}",
|
||||
"exit_code": 1,
|
||||
"untrusted_content": True,
|
||||
if remote:
|
||||
async with httpx.AsyncClient(timeout=15) as client:
|
||||
resp = await client.post(f"{_INTERNAL_BASE}/api/shell/exec",
|
||||
json={"command": cmd}, headers=headers)
|
||||
if resp.status_code >= 400:
|
||||
return {
|
||||
"error": f"shell/exec returned HTTP {resp.status_code}: {resp.text[:200]}",
|
||||
"exit_code": 1,
|
||||
"untrusted_content": True,
|
||||
}
|
||||
try:
|
||||
data = resp.json()
|
||||
except Exception:
|
||||
data = {}
|
||||
else:
|
||||
import asyncio
|
||||
proc = await asyncio.create_subprocess_exec(
|
||||
"tmux", "kill-session", "-t", session_id,
|
||||
stdout=asyncio.subprocess.PIPE,
|
||||
stderr=asyncio.subprocess.PIPE,
|
||||
)
|
||||
try:
|
||||
stdout, stderr = await asyncio.wait_for(proc.communicate(), timeout=5)
|
||||
except asyncio.TimeoutError:
|
||||
proc.kill()
|
||||
stdout, stderr = await proc.communicate()
|
||||
data = {
|
||||
"stdout": stdout.decode("utf-8", errors="replace"),
|
||||
"stderr": stderr.decode("utf-8", errors="replace"),
|
||||
"exit_code": proc.returncode,
|
||||
}
|
||||
try:
|
||||
data = resp.json()
|
||||
except Exception:
|
||||
data = {}
|
||||
kill_failed = isinstance(data, dict) and data.get("exit_code") not in (None, 0)
|
||||
kill_err = ((data.get("stderr") or data.get("error") or "").strip() if isinstance(data, dict) else "")
|
||||
# "no server running" / "can't find session" means it was already
|
||||
@@ -971,6 +1116,34 @@ async def _cookbook_kill_session(session_id: str, *, remote_host: str = "",
|
||||
if kill_failed and not already_gone:
|
||||
return {"error": f"Failed to {verb.lower()} {target_label}: {kill_err or 'kill-session returned non-zero'}", "exit_code": 1}
|
||||
|
||||
# Some model servers survive the tmux session's SIGHUP. For local
|
||||
# tracked tasks only, terminate processes whose full command line
|
||||
# exactly matches the command saved by the Cookbook launcher.
|
||||
if not remote and isinstance(matched, dict):
|
||||
import os
|
||||
import signal
|
||||
tracked_cmd = str((matched.get("payload") or {}).get("_cmd") or "").strip()
|
||||
matched_pids: list[int] = []
|
||||
if tracked_cmd:
|
||||
for pid_name in os.listdir("/proc"):
|
||||
if not pid_name.isdigit() or int(pid_name) == os.getpid():
|
||||
continue
|
||||
try:
|
||||
raw = open(f"/proc/{pid_name}/cmdline", "rb").read()
|
||||
process_cmd = raw.replace(b"\x00", b" ").decode("utf-8", errors="replace").strip()
|
||||
except (OSError, PermissionError):
|
||||
continue
|
||||
if process_cmd == tracked_cmd:
|
||||
matched_pids.append(int(pid_name))
|
||||
with contextlib.suppress(ProcessLookupError, PermissionError):
|
||||
os.kill(int(pid_name), signal.SIGTERM)
|
||||
if matched_pids:
|
||||
await asyncio.sleep(0.5)
|
||||
for pid in matched_pids:
|
||||
with contextlib.suppress(ProcessLookupError, PermissionError):
|
||||
os.kill(pid, 0)
|
||||
os.kill(pid, signal.SIGKILL)
|
||||
|
||||
# Update state: mark stopped (so the UI + list reflect reality).
|
||||
if matched is not None:
|
||||
try:
|
||||
@@ -1185,44 +1358,91 @@ async def do_cancel_download(content: str, owner: Optional[str] = None) -> Dict:
|
||||
|
||||
|
||||
async def do_search_hf_models(content: str, owner: Optional[str] = None) -> Dict:
|
||||
"""Search HuggingFace via the cookbook /api/cookbook/hf-latest endpoint."""
|
||||
from src.tool_implementations import _internal_headers, _INTERNAL_BASE # shared, lives in facade
|
||||
"""Search Hugging Face Hub models via the public HF API.
|
||||
|
||||
This intentionally does not use the cookbook's ``/hf-latest`` route:
|
||||
that route is a VRAM/trending browser and ignores semantic search terms.
|
||||
"""
|
||||
import httpx
|
||||
try:
|
||||
args = _parse_tool_args(content)
|
||||
except ValueError:
|
||||
return {"error": "Invalid JSON arguments", "exit_code": 1}
|
||||
query = args.get("query", "") or args.get("search", "")
|
||||
limit = args.get("limit", 10)
|
||||
params: Dict[str, str] = {}
|
||||
if query:
|
||||
query = _string_arg(args.get("query") or args.get("search") or args.get("q"))
|
||||
try:
|
||||
limit = max(1, min(int(args.get("limit") or 10), 25))
|
||||
except Exception:
|
||||
limit = 10
|
||||
explicit_author = _string_arg(args.get("author") or args.get("owner") or args.get("namespace"))
|
||||
official_only = bool(args.get("official_only") or args.get("official") or args.get("provider_only"))
|
||||
if re.search(r"\bofficial\b|\blatest\b|\bnewest\b|\brecent\b", query, flags=re.I):
|
||||
official_only = True
|
||||
official_author = explicit_author or (_hf_official_author_for_query(query) if official_only else "")
|
||||
wants_latest = bool(re.search(r"\blatest\b|\bnewest\b|\brecent\b|\breleased\b", query, flags=re.I))
|
||||
wants_quant = bool(args.get("quantized") or args.get("quant") or _hf_query_mentions_quant(query))
|
||||
exact_repo_query = _cookbook_is_exact_repo_id(query)
|
||||
if wants_quant and not (args.get("official_only") or args.get("official") or explicit_author):
|
||||
official_only = False
|
||||
official_author = ""
|
||||
|
||||
params: Dict[str, str] = {
|
||||
"limit": str(max(limit * 8, 50) if official_author else max(limit * 4, limit)),
|
||||
"full": "false",
|
||||
"sort": "lastModified" if wants_latest else "downloads",
|
||||
"direction": "-1",
|
||||
}
|
||||
if official_author:
|
||||
params["author"] = official_author
|
||||
elif query:
|
||||
params["search"] = query
|
||||
if limit:
|
||||
params["limit"] = str(limit)
|
||||
pipeline = _string_arg(args.get("pipeline") or args.get("filter"))
|
||||
if pipeline:
|
||||
params["filter"] = pipeline
|
||||
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=30) as client:
|
||||
resp = await client.get(f"{_INTERNAL_BASE}/api/cookbook/hf-latest",
|
||||
params=params, headers=_internal_headers())
|
||||
data = resp.json()
|
||||
models = data.get("models") if isinstance(data, dict) else data
|
||||
if not models:
|
||||
return {"output": f"No models found for query: {query!r}", "exit_code": 0}
|
||||
lines = [f"Found {len(models)} model(s) for {query!r}:" if query else f"{len(models)} model(s):"]
|
||||
for m in models[:limit if isinstance(limit, int) else 10]:
|
||||
if isinstance(m, dict):
|
||||
name = m.get("repo_id") or m.get("modelId") or m.get("id") or "?"
|
||||
dl = m.get("downloads")
|
||||
size = m.get("size_gb") or m.get("needed_vram_gb")
|
||||
bits = []
|
||||
if size:
|
||||
bits.append(f"~{size}GB")
|
||||
if dl:
|
||||
bits.append(f"{dl} downloads")
|
||||
tail = f" ({', '.join(bits)})" if bits else ""
|
||||
lines.append(f"- {name}{tail}")
|
||||
if exact_repo_query:
|
||||
resp = await client.get(f"https://huggingface.co/api/models/{query}")
|
||||
else:
|
||||
lines.append(f"- {m}")
|
||||
return {"output": "\n".join(lines), "models": models, "exit_code": 0}
|
||||
resp = await client.get("https://huggingface.co/api/models", params=params)
|
||||
if resp.status_code != 200:
|
||||
return {"error": f"HF API HTTP {resp.status_code}: {resp.text[:300]}", "exit_code": 1}
|
||||
data = resp.json()
|
||||
if isinstance(data, dict) and (data.get("id") or data.get("modelId")):
|
||||
models = [data]
|
||||
else:
|
||||
models = data if isinstance(data, list) else []
|
||||
if official_author:
|
||||
author_lc = official_author.lower()
|
||||
terms = _hf_query_terms(query, official_author)
|
||||
filtered = [
|
||||
m for m in models if isinstance(m, dict)
|
||||
and str(m.get("id") or m.get("modelId") or "").lower().startswith(f"{author_lc}/")
|
||||
and _hf_model_matches_terms(m, terms)
|
||||
]
|
||||
# For "latest official Qwen model", family/org is the only useful
|
||||
# constraint. If local term filtering removes every result, show
|
||||
# the official author's recent models rather than unrelated Hub hits.
|
||||
models = filtered or [
|
||||
m for m in models if isinstance(m, dict)
|
||||
and str(m.get("id") or m.get("modelId") or "").lower().startswith(f"{author_lc}/")
|
||||
]
|
||||
if not wants_quant and not exact_repo_query:
|
||||
models = [m for m in models if not _hf_model_is_quant_variant(m)]
|
||||
else:
|
||||
models = [m for m in models if isinstance(m, dict)]
|
||||
if not wants_quant:
|
||||
models = [m for m in models if not _hf_model_is_quant_variant(m)]
|
||||
models = models[:limit]
|
||||
if not models:
|
||||
scope = f" official author {official_author!r}" if official_author else ""
|
||||
return {"output": f"No{scope} models found for query: {query!r}", "models": [], "exit_code": 0}
|
||||
return {
|
||||
"output": _hf_format_model_search_output(models, query, official_author),
|
||||
"models": models,
|
||||
"official_author": official_author or None,
|
||||
"exit_code": 0,
|
||||
}
|
||||
except Exception as e:
|
||||
return {"error": str(e), "exit_code": 1}
|
||||
|
||||
@@ -1256,10 +1476,27 @@ async def do_adopt_served_model(content: str, owner: Optional[str] = None) -> Di
|
||||
port = args.get("port") or 8000
|
||||
display_name = (args.get("name") or "").strip() or (model.split("/")[-1] if "/" in model else model)
|
||||
add_endpoint = args.get("add_endpoint", True)
|
||||
dry_run = bool(args.get("dry_run", False))
|
||||
|
||||
if not sess or not model:
|
||||
return {"error": "tmux_session and model are required", "exit_code": 1}
|
||||
|
||||
if dry_run:
|
||||
return {
|
||||
"output": (
|
||||
f"Dry run: would verify tmux session {sess!r} on {host or 'local'}, "
|
||||
f"register model {model!r} on port {int(port)}, and "
|
||||
f"{'add' if add_endpoint else 'not add'} a chat endpoint. No state was changed."
|
||||
),
|
||||
"dry_run": True,
|
||||
"host": host,
|
||||
"tmux_session": sess,
|
||||
"model": model,
|
||||
"port": int(port),
|
||||
"add_endpoint": bool(add_endpoint),
|
||||
"exit_code": 0,
|
||||
}
|
||||
|
||||
# Verify tmux session exists on the target host
|
||||
if host:
|
||||
try:
|
||||
@@ -1460,6 +1697,7 @@ async def do_serve_preset(content: str, owner: Optional[str] = None) -> Dict:
|
||||
except ValueError:
|
||||
return {"error": "Invalid JSON arguments", "exit_code": 1}
|
||||
name = (args.get("name") or args.get("preset") or "").strip()
|
||||
dry_run = bool(args.get("dry_run", False))
|
||||
if not name:
|
||||
return {"error": "name (preset name) is required. Call list_serve_presets to see what's available.", "exit_code": 1}
|
||||
|
||||
@@ -1494,6 +1732,20 @@ async def do_serve_preset(content: str, owner: Optional[str] = None) -> Dict:
|
||||
if not repo_id or not cmd:
|
||||
return {"error": f"Preset {chosen.get('name')!r} is missing model or cmd — can't launch.", "exit_code": 1}
|
||||
|
||||
if dry_run:
|
||||
return {
|
||||
"output": (
|
||||
f"Dry run: would launch preset {chosen.get('name')!r}: {repo_id} "
|
||||
f"on {host or 'local'} with command {cmd!r}. No server was started."
|
||||
),
|
||||
"dry_run": True,
|
||||
"preset": chosen.get("name") or name,
|
||||
"model": repo_id,
|
||||
"host": host,
|
||||
"command": cmd,
|
||||
"exit_code": 0,
|
||||
}
|
||||
|
||||
payload: Dict[str, Any] = {"repo_id": repo_id, "cmd": cmd}
|
||||
if host:
|
||||
payload["remote_host"] = host
|
||||
@@ -1549,6 +1801,7 @@ async def do_list_cached_models(content: str, owner: Optional[str] = None) -> Di
|
||||
return {"error": "Invalid JSON arguments", "exit_code": 1}
|
||||
raw_host = (args.get("host") or "").strip()
|
||||
headers = _internal_headers()
|
||||
scan_errors = []
|
||||
|
||||
async def _scan_one(host_label: str, host_val: str, ssh_port: str = "",
|
||||
platform: str = "", model_dir: str = "") -> list:
|
||||
@@ -1573,13 +1826,18 @@ async def do_list_cached_models(content: str, owner: Optional[str] = None) -> Di
|
||||
async with httpx.AsyncClient(timeout=60) as client:
|
||||
resp = await client.get(f"{_INTERNAL_BASE}/api/model/cached",
|
||||
params=p, headers=headers)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
if isinstance(data, dict) and data.get('error'):
|
||||
raise ValueError('cache endpoint reported an error')
|
||||
ms = data.get("models", []) if isinstance(data, dict) else (data or [])
|
||||
for m in ms:
|
||||
m["host"] = host_label or "local"
|
||||
return ms or []
|
||||
except Exception as e:
|
||||
logger.debug(f"list_cached_models scan({host_label}) failed: {e}")
|
||||
status = getattr(getattr(e, 'response', None), 'status_code', None)
|
||||
scan_errors.append({'host': host_label or 'local', 'reason': f'HTTP {status}' if status else type(e).__name__})
|
||||
return []
|
||||
|
||||
# When the caller specifies a host explicitly, scan only that one (old behaviour).
|
||||
@@ -1592,10 +1850,13 @@ async def do_list_cached_models(content: str, owner: Optional[str] = None) -> Di
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=10) as client:
|
||||
st = await client.get(f"{_INTERNAL_BASE}/api/cookbook/state", headers=headers)
|
||||
st.raise_for_status()
|
||||
st_data = st.json() if st.headers.get("content-type", "").startswith("application/json") else {}
|
||||
servers = (st_data.get("env", {}) or {}).get("servers") or []
|
||||
except Exception as e:
|
||||
logger.debug(f"server list fetch failed: {e}")
|
||||
status = getattr(getattr(e, 'response', None), 'status_code', None)
|
||||
scan_errors.append({'host': 'server inventory', 'reason': f'HTTP {status}' if status else type(e).__name__})
|
||||
st_data = {}
|
||||
|
||||
def _dirs_for(server_record: Dict[str, Any]) -> str:
|
||||
@@ -1654,6 +1915,9 @@ async def do_list_cached_models(content: str, owner: Optional[str] = None) -> Di
|
||||
continue
|
||||
seen.add(key)
|
||||
models.append(m)
|
||||
if not models and scan_errors:
|
||||
return {'error': 'Cache inventory could not be verified; one or more server scans failed.',
|
||||
'scan_errors': scan_errors, 'models': [], 'exit_code': 1}
|
||||
if not models:
|
||||
# Cache scans can miss models downloaded into the HF default cache
|
||||
# when the server has no explicit model_dir configured. Surface
|
||||
@@ -1708,6 +1972,11 @@ async def do_list_cached_models(content: str, owner: Optional[str] = None) -> Di
|
||||
kind = " [diffusion]" if m.get("is_diffusion") else ""
|
||||
backend = f" ({m.get('backend')})" if m.get("backend") else ""
|
||||
lines.append(f"- {name}{kind}{backend} — {sz}{inc}")
|
||||
if scan_errors:
|
||||
warning = 'Cache inventory is incomplete; failed scans: ' + ', '.join(
|
||||
f"{item['host']} ({item['reason']})" for item in scan_errors)
|
||||
return {'output': warning + '\n\n' + '\n'.join(lines), 'models': models,
|
||||
'error': warning, 'scan_errors': scan_errors, 'partial': True, 'exit_code': 1}
|
||||
return {"output": "\n".join(lines), "models": models, "exit_code": 0}
|
||||
except Exception as e:
|
||||
return {"error": str(e), "exit_code": 1}
|
||||
|
||||
+98
-40
@@ -6,15 +6,17 @@ Holds the edit_image (gallery) tool.
|
||||
``_INTERNAL_BASE`` still lives in tool_implementations.py and is pulled back
|
||||
function-locally here.
|
||||
"""
|
||||
import hashlib
|
||||
import io
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
from typing import Dict, Optional
|
||||
|
||||
from src.tools._common import _parse_tool_args
|
||||
|
||||
|
||||
async def do_edit_image(content: str, owner: Optional[str] = None) -> Dict:
|
||||
"""Edit a gallery image (upscale, rembg, inpaint, harmonize)."""
|
||||
import httpx
|
||||
from src.tool_implementations import _INTERNAL_BASE # shared constant, still lives in the facade
|
||||
"""Create an owner-scoped edited copy of a gallery image."""
|
||||
try:
|
||||
args = _parse_tool_args(content)
|
||||
except ValueError:
|
||||
@@ -23,44 +25,100 @@ async def do_edit_image(content: str, owner: Optional[str] = None) -> Dict:
|
||||
action = args.get("action", "")
|
||||
if not image_id or not action:
|
||||
return {"error": "image_id and action are required", "exit_code": 1}
|
||||
payload = {"image_id": image_id}
|
||||
if args.get("prompt"):
|
||||
payload["prompt"] = args["prompt"]
|
||||
if args.get("scale"):
|
||||
payload["scale"] = args["scale"]
|
||||
if action not in {"upscale", "rembg"}:
|
||||
return {
|
||||
"error": f"Unsupported edit action: {action}. Use upscale or rembg.",
|
||||
"exit_code": 1,
|
||||
}
|
||||
|
||||
from core.database import GalleryImage, SessionLocal
|
||||
from src.constants import GENERATED_IMAGES_DIR
|
||||
|
||||
db = SessionLocal()
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=120) as client:
|
||||
resp = await client.post(f"{_INTERNAL_BASE}/api/gallery/{action}", json=payload)
|
||||
data = resp.json()
|
||||
new_id = data.get("id") or data.get("image_id")
|
||||
if data.get("success") or new_id:
|
||||
result = {
|
||||
"output": f"Image edited ({action}). New image ID: {new_id or '?'}",
|
||||
"exit_code": 0,
|
||||
}
|
||||
if new_id:
|
||||
result["image_id"] = new_id
|
||||
q = db.query(GalleryImage).filter(
|
||||
GalleryImage.id == image_id,
|
||||
GalleryImage.is_active == True, # noqa: E712
|
||||
)
|
||||
# A tool call without an owner must never fall through to another
|
||||
# user's gallery row.
|
||||
q = q.filter(GalleryImage.owner == owner) if owner else q.filter(False)
|
||||
source = q.first()
|
||||
if not source:
|
||||
return {"error": "Image not found", "exit_code": 1}
|
||||
|
||||
root = Path(GENERATED_IMAGES_DIR).resolve()
|
||||
source_name = Path(str(source.filename or "")).name
|
||||
source_path = (root / source_name).resolve()
|
||||
if source_name != source.filename or source_path.parent != root or not source_path.is_file():
|
||||
return {"error": "Image file not found", "exit_code": 1}
|
||||
|
||||
from PIL import Image
|
||||
|
||||
with Image.open(source_path) as opened:
|
||||
image = opened.convert("RGBA")
|
||||
if action == "upscale":
|
||||
try:
|
||||
from src.database import GalleryImage, SessionLocal
|
||||
db = SessionLocal()
|
||||
try:
|
||||
q = db.query(GalleryImage).filter(GalleryImage.id == new_id)
|
||||
if owner:
|
||||
q = q.filter(GalleryImage.owner == owner)
|
||||
img = q.first()
|
||||
if img and img.filename:
|
||||
result.update({
|
||||
"image_url": f"/api/generated-image/{img.filename}",
|
||||
"image_prompt": img.prompt or args.get("prompt") or action,
|
||||
"image_model": img.model or "edit_image",
|
||||
"image_size": img.size or "",
|
||||
"image_quality": img.quality or "",
|
||||
})
|
||||
finally:
|
||||
db.close()
|
||||
except Exception:
|
||||
pass
|
||||
return result
|
||||
return {"error": data.get("error", f"{action} failed"), "exit_code": 1}
|
||||
scale = int(args.get("scale") or 2)
|
||||
except (TypeError, ValueError):
|
||||
scale = 2
|
||||
if scale not in {2, 4}:
|
||||
return {"error": "scale must be 2 or 4", "exit_code": 1}
|
||||
image = image.resize(
|
||||
(image.width * scale, image.height * scale),
|
||||
Image.Resampling.LANCZOS,
|
||||
)
|
||||
else:
|
||||
try:
|
||||
from rembg import remove
|
||||
except ImportError:
|
||||
return {
|
||||
"error": "Background removal is not installed. Install the rembg optional dependency.",
|
||||
"exit_code": 1,
|
||||
}
|
||||
image = remove(image)
|
||||
|
||||
output = io.BytesIO()
|
||||
image.save(output, format="PNG")
|
||||
output_bytes = output.getvalue()
|
||||
width, height = image.size
|
||||
|
||||
root.mkdir(parents=True, exist_ok=True)
|
||||
filename = f"{uuid.uuid4().hex[:12]}.png"
|
||||
(root / filename).write_bytes(output_bytes)
|
||||
new_id = str(uuid.uuid4())
|
||||
derived = GalleryImage(
|
||||
id=new_id,
|
||||
filename=filename,
|
||||
prompt=source.prompt or action,
|
||||
caption=source.caption,
|
||||
model=f"edit_image:{action}",
|
||||
size=f"{width}x{height}",
|
||||
quality=source.quality,
|
||||
tags=source.tags,
|
||||
ai_tags=source.ai_tags,
|
||||
session_id=source.session_id,
|
||||
album_id=source.album_id,
|
||||
owner=owner,
|
||||
file_hash=hashlib.sha256(output_bytes).hexdigest(),
|
||||
file_size=len(output_bytes),
|
||||
width=width,
|
||||
height=height,
|
||||
)
|
||||
db.add(derived)
|
||||
db.commit()
|
||||
return {
|
||||
"output": f"Image edited ({action}). New image ID: {new_id}",
|
||||
"exit_code": 0,
|
||||
"image_id": new_id,
|
||||
"image_url": f"/api/generated-image/{filename}",
|
||||
"image_prompt": derived.prompt,
|
||||
"image_model": derived.model,
|
||||
"image_size": derived.size,
|
||||
"image_quality": derived.quality or "",
|
||||
}
|
||||
except Exception as e:
|
||||
db.rollback()
|
||||
return {"error": str(e), "exit_code": 1}
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
+196
-16
@@ -37,9 +37,24 @@ async def do_manage_notes(content: str, owner: Optional[str] = None) -> Dict:
|
||||
"save": "add",
|
||||
"remind": "add",
|
||||
"remove": "delete",
|
||||
"remove_item": "toggle_item",
|
||||
}
|
||||
action = _NOTE_ACTION_ALIASES.get(action, action)
|
||||
if action == "remove_item":
|
||||
return {
|
||||
"error": "To remove a checklist item, use update with id and the complete remaining checklist_items, preserving their done states. No item was changed.",
|
||||
"exit_code": 1,
|
||||
}
|
||||
list_search_query = str(
|
||||
args.get("search")
|
||||
or args.get("query")
|
||||
or args.get("text")
|
||||
or args.get("title")
|
||||
or args.get("content")
|
||||
or ""
|
||||
).strip()
|
||||
if action == "list" and list_search_query:
|
||||
action = "search"
|
||||
args.setdefault("query", list_search_query)
|
||||
db = SessionLocal()
|
||||
|
||||
def _norm_note_title(value: str) -> str:
|
||||
@@ -55,6 +70,9 @@ async def do_manage_notes(content: str, owner: Optional[str] = None) -> Dict:
|
||||
return True
|
||||
return getattr(note, "owner", None) == owner_value
|
||||
|
||||
def _is_calendar_reminder_note(note) -> bool:
|
||||
return getattr(note, "source", None) == "calendar" and getattr(note, "label", None) == "calendar"
|
||||
|
||||
def _note_by_prefix(note_id: str):
|
||||
if not note_id:
|
||||
return None
|
||||
@@ -63,15 +81,89 @@ async def do_manage_notes(content: str, owner: Optional[str] = None) -> Dict:
|
||||
q = q.filter(Note.owner == owner)
|
||||
return q.first()
|
||||
|
||||
def _format_note_list(notes) -> str:
|
||||
def _note_id_arg() -> str:
|
||||
return str(args.get("id") or args.get("note_id") or args.get("noteId") or "").strip()
|
||||
|
||||
def _norm_note_text(value) -> str:
|
||||
return re.sub(r"\s+", " ", str(value or "").strip())
|
||||
|
||||
def _norm_note_items(value) -> list[dict]:
|
||||
if value in (None, ""):
|
||||
return []
|
||||
raw = value
|
||||
if isinstance(raw, str):
|
||||
try:
|
||||
raw = json.loads(raw)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
return [{"text": _norm_note_text(raw), "done": False}]
|
||||
if not isinstance(raw, list):
|
||||
return [{"text": _norm_note_text(raw), "done": False}]
|
||||
items = []
|
||||
for item in raw:
|
||||
if isinstance(item, dict):
|
||||
text = _norm_note_text(item.get("text") or item.get("label") or item.get("title") or "")
|
||||
done = bool(item.get("done") or item.get("checked") or item.get("complete"))
|
||||
else:
|
||||
text = _norm_note_text(item)
|
||||
done = False
|
||||
if text:
|
||||
items.append({"text": text, "done": done})
|
||||
return items
|
||||
|
||||
def _existing_exact_note(
|
||||
*,
|
||||
title: str,
|
||||
content_value,
|
||||
items_value,
|
||||
note_type: str,
|
||||
label,
|
||||
due_date,
|
||||
color,
|
||||
pinned,
|
||||
):
|
||||
q = db.query(Note).filter(Note.archived == False) # noqa: E712
|
||||
if owner is not None:
|
||||
q = q.filter(Note.owner == owner)
|
||||
target_title = _norm_note_title(title)
|
||||
target_content = _norm_note_text(content_value)
|
||||
target_items = _norm_note_items(items_value)
|
||||
target_label = label or None
|
||||
target_due = due_date or None
|
||||
target_color = color or None
|
||||
target_pinned = bool(pinned)
|
||||
for existing in q.limit(50).all():
|
||||
if _norm_note_title(existing.title or "") != target_title:
|
||||
continue
|
||||
if (existing.note_type or "note") != (note_type or "note"):
|
||||
continue
|
||||
if (existing.label or None) != target_label:
|
||||
continue
|
||||
if (existing.due_date or None) != target_due:
|
||||
continue
|
||||
if (existing.color or None) != target_color:
|
||||
continue
|
||||
if bool(existing.pinned) != target_pinned:
|
||||
continue
|
||||
if _norm_note_text(existing.content) != target_content:
|
||||
continue
|
||||
if _norm_note_items(existing.items) != target_items:
|
||||
continue
|
||||
return existing
|
||||
return None
|
||||
|
||||
def _format_note_list(notes, *, full_content: bool = False) -> str:
|
||||
lines = []
|
||||
for n in notes:
|
||||
pin = " [PINNED]" if n.pinned else ""
|
||||
typ = " [checklist]" if n.note_type == "checklist" else ""
|
||||
lbl = f" #{n.label}" if n.label else ""
|
||||
title = n.title or "(untitled)"
|
||||
lines.append(f"- [{n.id[:8]}] **{title}**{pin}{typ}{lbl}")
|
||||
if n.note_type == "checklist" and n.items:
|
||||
lines.append(f"- [{n.id}] **{title}**{pin}{typ}{lbl}")
|
||||
# Search/list is a locator operation. Keep the body behind view so
|
||||
# a model cannot satisfy an explicit read request without the
|
||||
# required second call, and so large checklists do not flood the
|
||||
# next model context.
|
||||
if full_content and n.note_type == "checklist" and n.items:
|
||||
try:
|
||||
items = json.loads(n.items)
|
||||
for i, item in enumerate(items):
|
||||
@@ -79,9 +171,8 @@ async def do_manage_notes(content: str, owner: Optional[str] = None) -> Dict:
|
||||
lines.append(f" [{mark}] {i}: {item.get('text', '')}")
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
pass
|
||||
elif n.content:
|
||||
snippet = n.content[:80].replace("\n", " ")
|
||||
lines.append(f" {snippet}")
|
||||
elif full_content and n.content:
|
||||
lines.append(f" {n.content}")
|
||||
return "\n".join(lines)
|
||||
|
||||
try:
|
||||
@@ -95,6 +186,17 @@ async def do_manage_notes(content: str, owner: Optional[str] = None) -> Dict:
|
||||
show_archived = args.get("archived", False)
|
||||
q = q.filter(Note.archived == show_archived)
|
||||
notes = q.order_by(Note.pinned.desc(), Note.updated_at.desc()).all()
|
||||
if bool(args.get("pinned")):
|
||||
notes = [n for n in notes if bool(getattr(n, "pinned", False))]
|
||||
if bool(args.get("reminders") or args.get("due_only") or args.get("has_due_date")):
|
||||
notes = [n for n in notes if bool(getattr(n, "due_date", None))]
|
||||
include_calendar_reminders = bool(
|
||||
args.get("include_calendar_reminders")
|
||||
or str(args.get("source") or "").strip().lower() == "calendar"
|
||||
or str(args.get("label") or "").strip().lower() == "calendar-reminders"
|
||||
)
|
||||
if not include_calendar_reminders:
|
||||
notes = [n for n in notes if not _is_calendar_reminder_note(n)]
|
||||
if action in ("search", "find"):
|
||||
query = str(
|
||||
args.get("query")
|
||||
@@ -104,13 +206,20 @@ async def do_manage_notes(content: str, owner: Optional[str] = None) -> Dict:
|
||||
or ""
|
||||
).strip().lower()
|
||||
if query:
|
||||
query_terms = [
|
||||
term
|
||||
for term in re.findall(r"[a-z0-9]+", query)
|
||||
if term not in {"the", "a", "an", "note", "notes", "checklist", "list", "todo", "todos"}
|
||||
]
|
||||
filtered = []
|
||||
for n in notes:
|
||||
haystack = " ".join(
|
||||
str(part or "")
|
||||
for part in (n.title, n.content, n.label, n.items)
|
||||
).lower()
|
||||
if query in haystack:
|
||||
if query in haystack or (
|
||||
query_terms and all(term in haystack for term in query_terms)
|
||||
):
|
||||
filtered.append(n)
|
||||
notes = filtered
|
||||
if not notes:
|
||||
@@ -118,13 +227,13 @@ async def do_manage_notes(content: str, owner: Optional[str] = None) -> Dict:
|
||||
return {"results": _format_note_list(notes), "exit_code": 0}
|
||||
|
||||
elif action == "view":
|
||||
note_id = args.get("id", "")
|
||||
note_id = _note_id_arg()
|
||||
note = _note_by_prefix(note_id)
|
||||
if not note:
|
||||
return {"error": f"Note '{note_id}' not found", "exit_code": 1}
|
||||
if not _note_visible_to_owner(note, owner):
|
||||
return {"error": "Note not found", "exit_code": 1}
|
||||
return {"results": _format_note_list([note]), "exit_code": 0}
|
||||
return {"results": _format_note_list([note], full_content=True), "exit_code": 0}
|
||||
|
||||
elif action == "add":
|
||||
# Accept the various field names models emit: `text` is the most
|
||||
@@ -200,6 +309,25 @@ async def do_manage_notes(content: str, owner: Optional[str] = None) -> Dict:
|
||||
"duplicate": True,
|
||||
"exit_code": 0,
|
||||
}
|
||||
duplicate = _existing_exact_note(
|
||||
title=title,
|
||||
content_value=content_raw,
|
||||
items_value=items_raw,
|
||||
note_type=note_type,
|
||||
label=args.get("label"),
|
||||
due_date=due_iso,
|
||||
color=args.get("color"),
|
||||
pinned=args.get("pinned", False),
|
||||
)
|
||||
if duplicate:
|
||||
return {
|
||||
"response": f"Note already exists: \"{duplicate.title or title or '(untitled)'}\" (id: {duplicate.id[:8]})",
|
||||
"note_id": duplicate.id,
|
||||
"note_title": duplicate.title or title or "",
|
||||
"open_url": f"/#open=notes¬e={duplicate.id}",
|
||||
"duplicate": True,
|
||||
"exit_code": 0,
|
||||
}
|
||||
missing_id = reserve_upload_references(
|
||||
get_upload_handler(),
|
||||
owner,
|
||||
@@ -243,10 +371,33 @@ async def do_manage_notes(content: str, owner: Optional[str] = None) -> Dict:
|
||||
}
|
||||
|
||||
elif action == "update":
|
||||
note_id = args.get("id", "")
|
||||
note_id = _note_id_arg()
|
||||
note = _note_by_prefix(note_id)
|
||||
if not note:
|
||||
return {"error": f"Note '{note_id}' not found", "exit_code": 1}
|
||||
title_query = str(
|
||||
args.get("title")
|
||||
or args.get("query")
|
||||
or args.get("text")
|
||||
or ""
|
||||
).strip()
|
||||
if title_query:
|
||||
q = db.query(Note)
|
||||
if owner:
|
||||
q = q.filter(Note.owner == owner)
|
||||
candidates = [
|
||||
n for n in q.filter(Note.archived == False).all()
|
||||
if _norm_note_title(n.title) == _norm_note_title(title_query)
|
||||
]
|
||||
if len(candidates) == 1:
|
||||
note = candidates[0]
|
||||
elif len(candidates) > 1:
|
||||
return {
|
||||
"error": f"Multiple notes titled '{title_query}' found; pass an id.",
|
||||
"exit_code": 1,
|
||||
}
|
||||
if not note:
|
||||
target = note_id or args.get("title") or args.get("query") or args.get("text") or ""
|
||||
return {"error": f"Note '{target}' not found", "exit_code": 1}
|
||||
if not _note_visible_to_owner(note, owner):
|
||||
return {"error": "Note not found", "exit_code": 1}
|
||||
missing_id = reserve_upload_references(
|
||||
@@ -292,19 +443,46 @@ async def do_manage_notes(content: str, owner: Optional[str] = None) -> Dict:
|
||||
return {"response": f"Note updated: \"{note.title or '(untitled)'}\"", "exit_code": 0}
|
||||
|
||||
elif action == "delete":
|
||||
note_id = args.get("id", "")
|
||||
note_id = _note_id_arg()
|
||||
note = _note_by_prefix(note_id)
|
||||
if not note:
|
||||
return {"error": f"Note '{note_id}' not found", "exit_code": 1}
|
||||
title_query = str(
|
||||
args.get("title")
|
||||
or args.get("query")
|
||||
or args.get("text")
|
||||
or ""
|
||||
).strip()
|
||||
if title_query:
|
||||
q = db.query(Note)
|
||||
if owner:
|
||||
q = q.filter(Note.owner == owner)
|
||||
candidates = [
|
||||
n for n in q.filter(Note.archived == False).all()
|
||||
if _norm_note_title(n.title) == _norm_note_title(title_query)
|
||||
]
|
||||
if len(candidates) == 1:
|
||||
note = candidates[0]
|
||||
elif len(candidates) > 1:
|
||||
return {
|
||||
"error": f"Multiple notes titled '{title_query}' found; pass an id.",
|
||||
"exit_code": 1,
|
||||
}
|
||||
if not note:
|
||||
target = note_id or args.get("title") or args.get("query") or args.get("text") or ""
|
||||
return {"error": f"Note '{target}' not found", "exit_code": 1}
|
||||
if not _note_visible_to_owner(note, owner):
|
||||
return {"error": "Note not found", "exit_code": 1}
|
||||
title = note.title
|
||||
from src.tool_routing_experiment import note_fixture_scope
|
||||
fixture_scope = note_fixture_scope.get()
|
||||
if fixture_scope is not None and note.id not in fixture_scope:
|
||||
return {"error": "Target is outside the disposable test fixtures; no change made.", "exit_code": 1}
|
||||
db.delete(note)
|
||||
db.commit()
|
||||
return {"response": f"Deleted note: \"{title or '(untitled)'}\"", "exit_code": 0}
|
||||
|
||||
elif action == "toggle_item":
|
||||
note_id = args.get("id", "")
|
||||
note_id = _note_id_arg()
|
||||
index = args.get("index", 0)
|
||||
note = _note_by_prefix(note_id)
|
||||
if not note:
|
||||
@@ -316,7 +494,9 @@ async def do_manage_notes(content: str, owner: Optional[str] = None) -> Dict:
|
||||
items = json.loads(note.items)
|
||||
if index < 0 or index >= len(items):
|
||||
return {"error": f"Item index {index} out of range (0-{len(items)-1})", "exit_code": 1}
|
||||
items[index]["done"] = not items[index].get("done", False)
|
||||
if "done" in args and not isinstance(args["done"], bool):
|
||||
return {"error": "done must be a boolean (true or false)", "exit_code": 1}
|
||||
items[index]["done"] = args["done"] if "done" in args else not items[index].get("done", False)
|
||||
note.items = json.dumps(items)
|
||||
flag_modified(note, "items")
|
||||
db.commit()
|
||||
|
||||
+16
-2
@@ -88,11 +88,18 @@ async def do_manage_research(content: str, owner: Optional[str] = None) -> Dict:
|
||||
items.sort(reverse=True)
|
||||
if not items:
|
||||
return {"output": "No research found in the library." + (f" (search: {search})" if search else ""), "exit_code": 0}
|
||||
rows = "\n".join(f"- [{q or '(untitled)'}](#research-{sid}) — {n} sources" for _, sid, q, n in items[:50])
|
||||
# Keep the UI anchor and the API read identifier distinct. The anchor has
|
||||
# the `research-` UI prefix, while action=read expects the underlying file
|
||||
# stem. Exposing the exact id prevents agents from guessing or retrying
|
||||
# alternate spellings after a list call.
|
||||
rows = "\n".join(
|
||||
f"- [{q or '(untitled)'}](#research-{sid}) — id: {sid} — {n} sources"
|
||||
for _, sid, q, n in items[:50]
|
||||
)
|
||||
return {"output": f"Research library ({len(items)} item{'s' if len(items) != 1 else ''}):\n{rows}", "exit_code": 0}
|
||||
|
||||
|
||||
async def do_trigger_research(content: str, owner: Optional[str] = None) -> Dict:
|
||||
async def do_trigger_research(content: str, owner: Optional[str] = None, *, chat_session_id: Optional[str] = None) -> Dict:
|
||||
"""Start a live deep-research job that appears in the Deep Research
|
||||
sidebar. Hits /api/research/start (the same path the sidebar's
|
||||
'Research' button uses) so the session is discoverable + streamable
|
||||
@@ -107,10 +114,17 @@ async def do_trigger_research(content: str, owner: Optional[str] = None) -> Dict
|
||||
if not topic:
|
||||
return {"error": "topic (or query) is required", "exit_code": 1}
|
||||
payload: Dict[str, Any] = {"query": topic}
|
||||
if chat_session_id:
|
||||
# The dispatcher supplies the origin, never model-authored arguments.
|
||||
payload.update(origin_chat_id=chat_session_id, max_rounds=2, max_time=120)
|
||||
# Optional knobs the research panel supports.
|
||||
if args.get("max_rounds") is not None:
|
||||
try: payload["max_rounds"] = int(args["max_rounds"])
|
||||
except (ValueError, TypeError): pass
|
||||
if chat_session_id and payload.get('max_rounds') not in (1, 2):
|
||||
# Explicit deeper/Auto requests also regain the panel's normal time
|
||||
# budget; do not promise more rounds while keeping the quick cap.
|
||||
payload.pop('max_time', None)
|
||||
if args.get("max_time") is not None:
|
||||
try: payload["max_time"] = int(args["max_time"])
|
||||
except (ValueError, TypeError): pass
|
||||
|
||||
+29
-1
@@ -10,7 +10,12 @@ from typing import Dict
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def do_search_chats(query: str, limit: int = 20, owner: str | None = None) -> Dict:
|
||||
async def do_search_chats(
|
||||
query: str,
|
||||
limit: int = 20,
|
||||
owner: str | None = None,
|
||||
exclude_session_id: str | None = None,
|
||||
) -> Dict:
|
||||
"""Search past session transcripts for the calling user's sessions only.
|
||||
|
||||
Without an owner filter this used to leak EVERY user's chat history
|
||||
@@ -23,6 +28,29 @@ async def do_search_chats(query: str, limit: int = 20, owner: str | None = None)
|
||||
from src.session_search import search_session_messages
|
||||
|
||||
results = search_session_messages(query, limit=limit, owner=owner)
|
||||
if exclude_session_id:
|
||||
results = [r for r in results if r.session_id != exclude_session_id]
|
||||
if not results:
|
||||
from src.session_search import search_session_titles
|
||||
|
||||
results = search_session_titles(query, limit=limit, owner=owner)
|
||||
if exclude_session_id:
|
||||
results = [r for r in results if r.session_id != exclude_session_id]
|
||||
# Native callers often append the requested answer detail to a topic
|
||||
# query (for example, "roaster repair scheduling Jules"). Search the
|
||||
# topic prefix once when the exact full-text query misses; this keeps
|
||||
# chat retrieval useful without broadening into unrelated sessions.
|
||||
if not results:
|
||||
words = [word for word in query.split() if word]
|
||||
for width in (4, 3):
|
||||
if len(words) <= width:
|
||||
continue
|
||||
prefix = " ".join(words[:width])
|
||||
results = search_session_messages(prefix, limit=limit, owner=owner)
|
||||
if exclude_session_id:
|
||||
results = [r for r in results if r.session_id != exclude_session_id]
|
||||
if results:
|
||||
break
|
||||
if not results:
|
||||
return {"results": f"No chats found matching \"{query}\"."}
|
||||
|
||||
|
||||
+136
-40
@@ -77,6 +77,8 @@ async def do_manage_skills(content: str, owner: Optional[str] = None) -> Dict:
|
||||
if action == "view":
|
||||
if not name:
|
||||
return {"error": "name is required for view", "exit_code": 1}
|
||||
if args.get("path"):
|
||||
return {"error": "view reads SKILL.md only; use view_ref with name and path to read a supporting file.", "exit_code": 1}
|
||||
md = sm.read_skill_md(name, owner=owner)
|
||||
if md is None:
|
||||
return {"error": f"Skill {name!r} not found", "exit_code": 1}
|
||||
@@ -104,17 +106,11 @@ async def do_manage_skills(content: str, owner: Optional[str] = None) -> Dict:
|
||||
proc = args.get("steps") or []
|
||||
if not proc and not args.get("body_extra") and not args.get("solution"):
|
||||
return {"error": "procedure (or solution body) is required", "exit_code": 1}
|
||||
# Same auto-publish gate as the extractor path — when the user
|
||||
# has auto_approve_skills on and the caller didn't pin an explicit
|
||||
# status, publish immediately. Audit later demotes/removes on fail.
|
||||
# Newly learned procedures are always audited before they are allowed
|
||||
# into chat context. The automatic audit promotes passing skills.
|
||||
_status_arg = args.get("status")
|
||||
if not _status_arg:
|
||||
try:
|
||||
from routes.prefs_routes import _load_for_user as _load_prefs
|
||||
_prefs = _load_prefs(owner) or {}
|
||||
_status_arg = "published" if _prefs.get("auto_approve_skills", True) else "draft"
|
||||
except Exception:
|
||||
_status_arg = "draft"
|
||||
_status_arg = "draft"
|
||||
entry = sm.add_skill(
|
||||
name=args.get("name"),
|
||||
description=(args.get("description") or args.get("title") or "").strip(),
|
||||
@@ -162,7 +158,17 @@ async def do_manage_skills(content: str, owner: Optional[str] = None) -> Dict:
|
||||
return {"error": "name is required for edit", "exit_code": 1}
|
||||
new_content = args.get("content")
|
||||
if not isinstance(new_content, str) or not new_content.strip():
|
||||
return {"error": "content (full SKILL.md) is required for edit", "exit_code": 1}
|
||||
metadata_updates = {
|
||||
key: args[key] for key in (
|
||||
"description", "category", "when_to_use", "version", "confidence",
|
||||
"tags", "platforms", "requires_toolsets", "fallback_for_toolsets",
|
||||
"procedure", "pitfalls", "verification",
|
||||
) if key in args
|
||||
}
|
||||
if not metadata_updates:
|
||||
return {"error": "content (full SKILL.md) or an editable metadata field is required for edit", "exit_code": 1}
|
||||
ok = sm.update_skill(name, metadata_updates, owner=owner)
|
||||
return {"results": f"Edited skill `{name}`."} if ok else {"error": "Skill not found or update failed", "exit_code": 1}
|
||||
try:
|
||||
sk_new = Skill.from_markdown(new_content)
|
||||
except Exception as e:
|
||||
@@ -184,6 +190,8 @@ async def do_manage_skills(content: str, owner: Optional[str] = None) -> Dict:
|
||||
new_str = args.get("new_string", "")
|
||||
if not isinstance(old, str) or not old:
|
||||
return {"error": "old_string is required and must be non-empty", "exit_code": 1}
|
||||
if not isinstance(new_str, str):
|
||||
return {"error": "new_string must be a string; use an empty string to remove text", "exit_code": 1}
|
||||
md = sm.read_skill_md(name, owner=owner)
|
||||
if md is None:
|
||||
return {"error": f"Skill {name!r} not found", "exit_code": 1}
|
||||
@@ -211,8 +219,9 @@ async def do_manage_skills(content: str, owner: Optional[str] = None) -> Dict:
|
||||
updates = {"status": "published"}
|
||||
if args.get("confidence") is not None:
|
||||
updates["confidence"] = max(0.0, min(1.0, float(args["confidence"])))
|
||||
sm.update_skill(name, updates, owner=owner)
|
||||
return {"results": f"✅ Published `{name}`. It now appears in the skills index for future turns."}
|
||||
if not sm.update_skill(name, updates, owner=owner):
|
||||
return {"error": "Skill could not be published; no update was saved.", "exit_code": 1}
|
||||
return {"results": f"Published `{name}`. Automatic use remains subject to skill audit and approval settings."}
|
||||
|
||||
if action == "delete":
|
||||
if not name:
|
||||
@@ -271,6 +280,19 @@ def _skill_dump(sk) -> Dict:
|
||||
# Task management tool
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _task_date_utc(value):
|
||||
"""Parse the advertised one-off ISO datetime into the DB's naive UTC."""
|
||||
from datetime import datetime, timezone
|
||||
if not isinstance(value, str) or not value.strip():
|
||||
raise ValueError("scheduled_date is required for a one-off task")
|
||||
try:
|
||||
parsed = datetime.fromisoformat(value.strip().replace('Z', '+00:00'))
|
||||
except ValueError as exc:
|
||||
raise ValueError("scheduled_date must be an ISO datetime") from exc
|
||||
if parsed.tzinfo is not None:
|
||||
parsed = parsed.astimezone(timezone.utc).replace(tzinfo=None)
|
||||
return parsed
|
||||
|
||||
async def do_manage_tasks(content: str, owner: Optional[str] = None) -> Dict:
|
||||
"""Handle manage_tasks tool calls: CRUD on scheduled tasks."""
|
||||
import uuid as _uuid
|
||||
@@ -308,13 +330,79 @@ async def do_manage_tasks(content: str, owner: Optional[str] = None) -> Dict:
|
||||
args["scheduled_day"] = days[day]
|
||||
db = SessionLocal()
|
||||
try:
|
||||
def _task_by_id_or_exact_name(required_for: str):
|
||||
task_id = args.get("task_id")
|
||||
if task_id:
|
||||
task = db.query(ScheduledTask).filter(ScheduledTask.id == task_id).first()
|
||||
if task:
|
||||
return task, None
|
||||
|
||||
try:
|
||||
_uuid.UUID(str(task_id))
|
||||
looks_like_uuid = True
|
||||
except (TypeError, ValueError):
|
||||
looks_like_uuid = False
|
||||
if looks_like_uuid:
|
||||
return None, {"error": f"Task {task_id} not found", "exit_code": 1}
|
||||
|
||||
q = db.query(ScheduledTask).filter(ScheduledTask.name == str(task_id).strip())
|
||||
if owner:
|
||||
q = q.filter(ScheduledTask.owner == owner)
|
||||
matches = q.order_by(ScheduledTask.created_at.desc()).all()
|
||||
if len(matches) == 1:
|
||||
return matches[0], None
|
||||
if len(matches) > 1:
|
||||
return None, {
|
||||
"error": f"Task name '{task_id}' matched {len(matches)} tasks; use task_id",
|
||||
"exit_code": 1,
|
||||
}
|
||||
return None, {"error": f"Task {task_id} not found", "exit_code": 1}
|
||||
|
||||
name = str(args.get("name") or "").strip()
|
||||
if not name:
|
||||
return None, {"error": f"task_id is required for {required_for}", "exit_code": 1}
|
||||
|
||||
q = db.query(ScheduledTask).filter(ScheduledTask.name == name)
|
||||
if owner:
|
||||
q = q.filter(ScheduledTask.owner == owner)
|
||||
matches = q.order_by(ScheduledTask.created_at.desc()).all()
|
||||
if not matches:
|
||||
return None, {"error": f"Task named '{name}' not found", "exit_code": 1}
|
||||
if len(matches) > 1:
|
||||
return None, {
|
||||
"error": f"Task name '{name}' matched {len(matches)} tasks; use task_id",
|
||||
"exit_code": 1,
|
||||
}
|
||||
return matches[0], None
|
||||
|
||||
if action == "list":
|
||||
q = db.query(ScheduledTask)
|
||||
if owner:
|
||||
q = q.filter(ScheduledTask.owner == owner)
|
||||
status_filter = str(args.get("status") or "").strip().lower()
|
||||
if status_filter:
|
||||
q = q.filter(ScheduledTask.status == status_filter)
|
||||
name_filter = str(args.get("name") or "").strip()
|
||||
query_filter = str(
|
||||
args.get("query")
|
||||
or args.get("search")
|
||||
or args.get("pattern")
|
||||
or args.get("prompt")
|
||||
or args.get("match")
|
||||
or ""
|
||||
).strip()
|
||||
if name_filter:
|
||||
q = q.filter(ScheduledTask.name == name_filter)
|
||||
elif query_filter:
|
||||
from sqlalchemy import or_
|
||||
q = q.filter(or_(
|
||||
ScheduledTask.name.contains(query_filter),
|
||||
ScheduledTask.prompt.contains(query_filter),
|
||||
))
|
||||
tasks = q.order_by(ScheduledTask.created_at.desc()).all()
|
||||
if not tasks:
|
||||
return {"response": "No scheduled tasks found.", "exit_code": 0}
|
||||
suffix = f" matching '{name_filter or query_filter}'" if (name_filter or query_filter) else ""
|
||||
return {"response": f"No scheduled tasks found{suffix}.", "exit_code": 0}
|
||||
|
||||
lines = [f"Found {len(tasks)} tasks:"]
|
||||
for idx, t in enumerate(tasks, 1):
|
||||
@@ -326,6 +414,8 @@ async def do_manage_tasks(content: str, owner: Optional[str] = None) -> Dict:
|
||||
if t.next_run:
|
||||
bits.append(f"next {t.next_run.isoformat()}Z")
|
||||
detail = ", ".join(bits)
|
||||
if t.prompt:
|
||||
detail = f"{detail}; prompt: {t.prompt}"
|
||||
lines.append(f"{idx}. {t.name} ({t.id}) — {detail}")
|
||||
return {"response": "\n".join(lines), "exit_code": 0}
|
||||
|
||||
@@ -340,12 +430,17 @@ async def do_manage_tasks(content: str, owner: Optional[str] = None) -> Dict:
|
||||
|
||||
# Compute next_run for schedule triggers
|
||||
next_run = None
|
||||
scheduled_date = None
|
||||
if trigger_type == "schedule":
|
||||
schedule = args.get("schedule", "daily")
|
||||
if schedule == "once":
|
||||
scheduled_date = _task_date_utc(args.get("scheduled_date"))
|
||||
next_run = compute_next_run(
|
||||
schedule, args.get("scheduled_time", "09:00"),
|
||||
args.get("scheduled_day"),
|
||||
args.get("scheduled_day"), scheduled_date,
|
||||
)
|
||||
if schedule == "once" and next_run is None:
|
||||
return {"error": "scheduled_date must be in the future", "exit_code": 1}
|
||||
|
||||
task_id = str(_uuid.uuid4())
|
||||
# Guard each fallback with `or`: args.get("prompt", default) returns
|
||||
@@ -359,9 +454,10 @@ async def do_manage_tasks(content: str, owner: Optional[str] = None) -> Dict:
|
||||
prompt=args.get("prompt"),
|
||||
task_type=task_type,
|
||||
action=args.get("action_name"),
|
||||
schedule=args.get("schedule") if trigger_type == "schedule" else None,
|
||||
schedule=args.get("schedule", "daily") if trigger_type == "schedule" else None,
|
||||
scheduled_time=args.get("scheduled_time", "09:00") if trigger_type == "schedule" else None,
|
||||
scheduled_day=args.get("scheduled_day"),
|
||||
scheduled_date=scheduled_date,
|
||||
trigger_type=trigger_type,
|
||||
trigger_event=args.get("trigger_event"),
|
||||
trigger_count=args.get("trigger_count"),
|
||||
@@ -375,12 +471,9 @@ async def do_manage_tasks(content: str, owner: Optional[str] = None) -> Dict:
|
||||
return {"response": f"Created task '{name}' (id: {task_id})", "task_id": task_id, "exit_code": 0}
|
||||
|
||||
elif action == "edit":
|
||||
task_id = args.get("task_id")
|
||||
if not task_id:
|
||||
return {"error": "task_id is required for edit", "exit_code": 1}
|
||||
task = db.query(ScheduledTask).filter(ScheduledTask.id == task_id).first()
|
||||
if not task:
|
||||
return {"error": f"Task {task_id} not found", "exit_code": 1}
|
||||
task, error = _task_by_id_or_exact_name("edit")
|
||||
if error:
|
||||
return error
|
||||
# Strict ownership: the old `task.owner and task.owner != owner`
|
||||
# skipped the check on an owner-less task (created in no-login mode
|
||||
# or before the legacy-owner sweep), letting any authenticated user
|
||||
@@ -415,22 +508,28 @@ async def do_manage_tasks(content: str, owner: Optional[str] = None) -> Dict:
|
||||
setattr(task, field, args[field])
|
||||
changed.append(field)
|
||||
schedule_changed = True
|
||||
if "scheduled_date" in args:
|
||||
task.scheduled_date = _task_date_utc(args["scheduled_date"])
|
||||
changed.append("scheduled_date")
|
||||
schedule_changed = True
|
||||
|
||||
if schedule_changed and (task.trigger_type or "schedule") == "schedule":
|
||||
if task.schedule == "once" and task.scheduled_date is None:
|
||||
raise ValueError("scheduled_date is required for a one-off task")
|
||||
task.next_run = compute_next_run(
|
||||
task.schedule, task.scheduled_time, task.scheduled_day,
|
||||
task.scheduled_date,
|
||||
)
|
||||
if task.schedule == "once" and task.next_run is None:
|
||||
raise ValueError("scheduled_date must be in the future")
|
||||
|
||||
db.commit()
|
||||
return {"response": f"Updated task '{task.name}': {', '.join(changed)}", "exit_code": 0}
|
||||
|
||||
elif action == "delete":
|
||||
task_id = args.get("task_id")
|
||||
if not task_id:
|
||||
return {"error": "task_id is required for delete", "exit_code": 1}
|
||||
task = db.query(ScheduledTask).filter(ScheduledTask.id == task_id).first()
|
||||
if not task:
|
||||
return {"error": f"Task {task_id} not found", "exit_code": 1}
|
||||
task, error = _task_by_id_or_exact_name("delete")
|
||||
if error:
|
||||
return error
|
||||
if owner and task.owner != owner:
|
||||
return {"error": "Access denied", "exit_code": 1}
|
||||
name = task.name
|
||||
@@ -439,12 +538,9 @@ async def do_manage_tasks(content: str, owner: Optional[str] = None) -> Dict:
|
||||
return {"response": f"Deleted task '{name}'", "exit_code": 0}
|
||||
|
||||
elif action in ("pause", "resume"):
|
||||
task_id = args.get("task_id")
|
||||
if not task_id:
|
||||
return {"error": f"task_id is required for {action}", "exit_code": 1}
|
||||
task = db.query(ScheduledTask).filter(ScheduledTask.id == task_id).first()
|
||||
if not task:
|
||||
return {"error": f"Task {task_id} not found", "exit_code": 1}
|
||||
task, error = _task_by_id_or_exact_name(action)
|
||||
if error:
|
||||
return error
|
||||
if owner and task.owner != owner:
|
||||
return {"error": "Access denied", "exit_code": 1}
|
||||
|
||||
@@ -455,24 +551,24 @@ async def do_manage_tasks(content: str, owner: Optional[str] = None) -> Dict:
|
||||
if (task.trigger_type or "schedule") == "schedule":
|
||||
task.next_run = compute_next_run(
|
||||
task.schedule, task.scheduled_time, task.scheduled_day,
|
||||
task.scheduled_date,
|
||||
)
|
||||
if task.schedule == "once" and task.next_run is None:
|
||||
raise ValueError("A future scheduled_date is required to resume this one-off task")
|
||||
db.commit()
|
||||
return {"response": f"Task '{task.name}' {action}d", "exit_code": 0}
|
||||
|
||||
elif action == "run":
|
||||
task_id = args.get("task_id")
|
||||
if not task_id:
|
||||
return {"error": "task_id is required for run", "exit_code": 1}
|
||||
task = db.query(ScheduledTask).filter(ScheduledTask.id == task_id).first()
|
||||
if not task:
|
||||
return {"error": f"Task {task_id} not found", "exit_code": 1}
|
||||
task, error = _task_by_id_or_exact_name("run")
|
||||
if error:
|
||||
return error
|
||||
if owner and task.owner != owner:
|
||||
return {"error": "Access denied", "exit_code": 1}
|
||||
|
||||
from src.event_bus import get_task_scheduler
|
||||
scheduler = get_task_scheduler()
|
||||
if scheduler:
|
||||
started = await scheduler.run_task_now(task_id)
|
||||
started = await scheduler.run_task_now(task.id)
|
||||
if started:
|
||||
return {"response": f"Task '{task.name}' triggered", "exit_code": 0}
|
||||
else:
|
||||
|
||||
Reference in New Issue
Block a user