Merge pull request #1 from djw4/copilot/steward-concept-research
This commit is contained in:
commit
50d482d697
37
.dockerignore
Normal file
37
.dockerignore
Normal file
@ -0,0 +1,37 @@
|
||||
# Python
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
*.egg-info/
|
||||
.eggs/
|
||||
dist/
|
||||
build/
|
||||
|
||||
# Virtual environments
|
||||
.venv/
|
||||
venv/
|
||||
env/
|
||||
|
||||
# Environment secrets — never bake these into the image
|
||||
.env
|
||||
|
||||
# Test / lint caches
|
||||
.pytest_cache/
|
||||
.coverage
|
||||
htmlcov/
|
||||
.mypy_cache/
|
||||
.ruff_cache/
|
||||
|
||||
# IDE
|
||||
.vscode/
|
||||
.idea/
|
||||
|
||||
# OS
|
||||
.DS_Store
|
||||
|
||||
# Git
|
||||
.git/
|
||||
.github/
|
||||
|
||||
# Data directory (runtime artefact, not part of the image)
|
||||
data/
|
||||
thread_memory.json
|
||||
32
.env.example
Normal file
32
.env.example
Normal file
@ -0,0 +1,32 @@
|
||||
# Copy this file to .env and fill in your values.
|
||||
# Never commit .env to version control.
|
||||
|
||||
# Telegram bot token from @BotFather
|
||||
TELEGRAM_BOT_TOKEN=
|
||||
|
||||
# Comma-separated Telegram user IDs allowed to talk to Steward.
|
||||
# Leave empty to allow everyone (not recommended for production).
|
||||
TELEGRAM_ALLOWED_USER_IDS=
|
||||
|
||||
# OpenAI (or compatible) credentials
|
||||
OPENAI_API_KEY=
|
||||
OPENAI_BASE_URL=https://api.openai.com/v1
|
||||
OPENAI_MODEL=gpt-4o
|
||||
|
||||
# Optional: override the default system prompt
|
||||
# OPENAI_SYSTEM_PROMPT=
|
||||
|
||||
# Daily analysis target (optional)
|
||||
# ANALYSIS_TARGET_URL=https://your-internal-api/endpoint
|
||||
# ANALYSIS_TARGET_API_KEY=
|
||||
# ANALYSIS_CRON_HOUR=8
|
||||
# ANALYSIS_CRON_MINUTE=0
|
||||
|
||||
# Thread memory store path (JSON file for persisted thread summaries)
|
||||
# THREAD_MEMORY_PATH=thread_memory.json
|
||||
|
||||
# MCP / OpenAPI tool server (open-webui/openapi-servers compatible)
|
||||
# Point to any OpenAPI-spec tool server to enable LLM tool calling.
|
||||
# The service fetches /openapi.json from MCP_SERVER_URL to discover tools.
|
||||
# MCP_SERVER_URL=http://localhost:8000
|
||||
# MCP_SERVER_API_KEY=
|
||||
21
.github/pull_request_template.md
vendored
Normal file
21
.github/pull_request_template.md
vendored
Normal file
@ -0,0 +1,21 @@
|
||||
## What does this PR do?
|
||||
|
||||
<!-- A brief description of the change and why it's needed. -->
|
||||
|
||||
## Type of change
|
||||
|
||||
- [ ] `fix:` Bug fix (patch release)
|
||||
- [ ] `feat:` New feature (minor release)
|
||||
- [ ] `feat!:` / `BREAKING CHANGE:` Breaking change (major release)
|
||||
- [ ] `docs:` Documentation only
|
||||
- [ ] `chore:` / `refactor:` / `test:` No release
|
||||
|
||||
## Notes
|
||||
|
||||
<!-- Anything reviewers should pay special attention to, or context that doesn't fit above. -->
|
||||
|
||||
---
|
||||
|
||||
> **Commit message format** – this project uses [Conventional Commits](https://www.conventionalcommits.org/)
|
||||
> to drive automatic semantic versioning via `release-please`.
|
||||
> Use `fix:`, `feat:`, or `feat!:` as your commit prefix so the version bump is picked up correctly.
|
||||
62
.github/workflows/ci.yml
vendored
Normal file
62
.github/workflows/ci.yml
vendored
Normal file
@ -0,0 +1,62 @@
|
||||
name: CI
|
||||
|
||||
on:
|
||||
push:
|
||||
branches-ignore:
|
||||
- main
|
||||
pull_request:
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
packages: write
|
||||
|
||||
jobs:
|
||||
lint-and-test:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
cache: pip
|
||||
|
||||
- run: pip install -e ".[dev]"
|
||||
- run: python -m ruff check steward/ tests/
|
||||
- run: python -m pytest tests/ -v
|
||||
|
||||
docker-build:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Log in to GHCR
|
||||
# Only log in when pushing (not on PRs from forks)
|
||||
if: github.event_name == 'push'
|
||||
uses: docker/login-action@v3
|
||||
with:
|
||||
registry: ghcr.io
|
||||
username: ${{ github.actor }}
|
||||
password: ${{ secrets.GITHUB_TOKEN }}
|
||||
|
||||
- name: Docker metadata
|
||||
id: meta
|
||||
uses: docker/metadata-action@v5
|
||||
with:
|
||||
images: ghcr.io/${{ github.repository }}
|
||||
tags: |
|
||||
# feature-branch push → ghcr.io/…/steward:my-feature-branch
|
||||
type=ref,event=branch
|
||||
# pull-request → ghcr.io/…/steward:pr-42
|
||||
type=ref,event=pr
|
||||
# always → ghcr.io/…/steward:sha-abc1234
|
||||
type=sha,prefix=sha-
|
||||
|
||||
- name: Build and push
|
||||
uses: docker/build-push-action@v6
|
||||
with:
|
||||
context: .
|
||||
# Push on branch pushes; only validate (no push) on PRs
|
||||
push: ${{ github.event_name == 'push' }}
|
||||
tags: ${{ steps.meta.outputs.tags }}
|
||||
labels: ${{ steps.meta.outputs.labels }}
|
||||
62
.github/workflows/release.yml
vendored
Normal file
62
.github/workflows/release.yml
vendored
Normal file
@ -0,0 +1,62 @@
|
||||
name: Release
|
||||
|
||||
# Runs on every push to main.
|
||||
# release-please analyses conventional commits since the last release and either:
|
||||
# • creates/updates a "Release PR" that bumps pyproject.toml + CHANGELOG.md, or
|
||||
# • (when that PR is merged) tags the repo and creates a GitHub Release.
|
||||
# Only when a GitHub Release is actually created does the publish job run.
|
||||
|
||||
on:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
|
||||
permissions:
|
||||
contents: write
|
||||
pull-requests: write
|
||||
packages: write
|
||||
|
||||
jobs:
|
||||
release-please:
|
||||
runs-on: ubuntu-latest
|
||||
outputs:
|
||||
release_created: ${{ steps.release.outputs.release_created }}
|
||||
tag_name: ${{ steps.release.outputs.tag_name }}
|
||||
steps:
|
||||
- uses: googleapis/release-please-action@v4
|
||||
id: release
|
||||
with:
|
||||
release-type: python
|
||||
|
||||
publish:
|
||||
needs: release-please
|
||||
if: needs.release-please.outputs.release_created == 'true'
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Log in to GHCR
|
||||
uses: docker/login-action@v3
|
||||
with:
|
||||
registry: ghcr.io
|
||||
username: ${{ github.actor }}
|
||||
password: ${{ secrets.GITHUB_TOKEN }}
|
||||
|
||||
- name: Docker metadata
|
||||
id: meta
|
||||
uses: docker/metadata-action@v5
|
||||
with:
|
||||
images: ghcr.io/${{ github.repository }}
|
||||
tags: |
|
||||
# semver tag → ghcr.io/…/steward:1.2.3
|
||||
type=raw,value=${{ needs.release-please.outputs.tag_name }}
|
||||
# always push latest on a real release
|
||||
type=raw,value=latest
|
||||
|
||||
- name: Build and push
|
||||
uses: docker/build-push-action@v6
|
||||
with:
|
||||
context: .
|
||||
push: true
|
||||
tags: ${{ steps.meta.outputs.tags }}
|
||||
labels: ${{ steps.meta.outputs.labels }}
|
||||
37
.gitignore
vendored
Normal file
37
.gitignore
vendored
Normal file
@ -0,0 +1,37 @@
|
||||
# Python
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
*.pyo
|
||||
*.pyd
|
||||
*.so
|
||||
*.egg
|
||||
*.egg-info/
|
||||
dist/
|
||||
build/
|
||||
.eggs/
|
||||
|
||||
# Virtual environments
|
||||
.venv/
|
||||
venv/
|
||||
env/
|
||||
|
||||
# Environment
|
||||
.env
|
||||
|
||||
# Pytest
|
||||
.pytest_cache/
|
||||
.coverage
|
||||
htmlcov/
|
||||
|
||||
# Mypy
|
||||
.mypy_cache/
|
||||
|
||||
# Ruff
|
||||
.ruff_cache/
|
||||
|
||||
# IDE
|
||||
.vscode/
|
||||
.idea/
|
||||
|
||||
# OS
|
||||
.DS_Store
|
||||
3
.release-please-manifest.json
Normal file
3
.release-please-manifest.json
Normal file
@ -0,0 +1,3 @@
|
||||
{
|
||||
".": "0.1.0"
|
||||
}
|
||||
23
Dockerfile
Normal file
23
Dockerfile
Normal file
@ -0,0 +1,23 @@
|
||||
FROM python:3.12-slim
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
# Install dependencies (separate layer for cache efficiency)
|
||||
COPY pyproject.toml README.md ./
|
||||
COPY steward/ steward/
|
||||
RUN pip install --no-cache-dir .
|
||||
|
||||
# Create a non-root user (UID/GID 1000) and persistent data directory
|
||||
RUN addgroup --gid 1000 steward \
|
||||
&& adduser --uid 1000 --gid 1000 --no-create-home --disabled-password --gecos "" steward \
|
||||
&& mkdir /data \
|
||||
&& chown steward:steward /data
|
||||
|
||||
USER steward
|
||||
|
||||
# Default path for the thread memory store — override via env or docker-compose
|
||||
ENV THREAD_MEMORY_PATH=/data/thread_memory.json
|
||||
|
||||
VOLUME ["/data"]
|
||||
|
||||
CMD ["steward"]
|
||||
516
NORTH_STAR.md
Normal file
516
NORTH_STAR.md
Normal file
@ -0,0 +1,516 @@
|
||||
# Steward
|
||||
## Vision and Intent
|
||||
|
||||
Version: 0.1 (Concept)
|
||||
Status: Research
|
||||
Author: Daniel Wagner
|
||||
|
||||
---
|
||||
|
||||
# Executive Summary
|
||||
|
||||
Steward is a long-running, AI-assisted personal operations platform designed to reduce cognitive load by acting as a persistent, trustworthy steward of both digital infrastructure and delegated personal objectives.
|
||||
|
||||
Unlike traditional AI assistants, Steward is not intended to answer prompts in isolation. Instead, it maintains an ongoing understanding of its environment, remembers historical context, continuously evaluates events, and takes appropriate action within explicitly delegated authority.
|
||||
|
||||
The primary research question underpinning Steward is:
|
||||
|
||||
> Can an AI system safely become more useful over time without becoming less predictable?
|
||||
|
||||
Steward seeks to answer this through carefully constrained autonomy, policy-driven execution and continuous learning from operational history.
|
||||
|
||||
---
|
||||
|
||||
# Philosophy
|
||||
|
||||
Steward is not *just* an AI chatbot.
|
||||
|
||||
Steward is not an autonomous root user.
|
||||
|
||||
Steward is not *just* an automation engine.
|
||||
|
||||
Steward is a persistent operations platform.
|
||||
|
||||
Its purpose is to quietly reduce cognitive load by:
|
||||
|
||||
- observing
|
||||
- remembering
|
||||
- planning
|
||||
- researching
|
||||
- proposing
|
||||
- executing safe actions
|
||||
- continuously improving
|
||||
|
||||
The system should feel less like ChatGPT and more like employing a careful, diligent junior platform engineer and executive assistant who never forgets anything.
|
||||
|
||||
---
|
||||
|
||||
# Core Principles
|
||||
|
||||
## Persistent
|
||||
|
||||
Steward has persistent state and is never "reset" between interactions.
|
||||
|
||||
Conceptually it is always running. In practice this means persistent state with event-driven wake, not a literal always-on compute loop, so persistence is not bought at the cost of continuous compute.
|
||||
|
||||
It continuously evaluates changes in its environment as they arrive.
|
||||
|
||||
Scheduled tasks are allowed but the initiator is an event.
|
||||
|
||||
---
|
||||
|
||||
## Event Driven
|
||||
|
||||
Everything is an event.
|
||||
|
||||
Examples:
|
||||
|
||||
- Prometheus alert
|
||||
- Telegram conversation
|
||||
- Calendar update
|
||||
- Todoist completion
|
||||
- GitHub Pull Request
|
||||
- AWX job completion
|
||||
- Weather change
|
||||
- Home Assistant state change
|
||||
|
||||
Events are evaluated against active goals.
|
||||
|
||||
---
|
||||
|
||||
## Event Processing
|
||||
|
||||
"Everything is an event" does not mean everything is evaluated.
|
||||
|
||||
A homelab emits thousands of low-value events, and metrics sources can flap continuously. Raw events pass through a processing layer before reaching goal evaluation:
|
||||
|
||||
- **Ingestion** normalises events from all sources into a common shape.
|
||||
- **Filtering** drops noise below a relevance threshold.
|
||||
- **Deduplication and debounce** collapse repeated or flapping events.
|
||||
- **Correlation** groups related events into a single *situation* (e.g. three alerts from one failing disk).
|
||||
|
||||
Steward evaluates *situations* against goals, not raw events. This bounds compute and prevents notification storms.
|
||||
|
||||
Event content is untrusted input. See [Security Model](#security-model).
|
||||
|
||||
---
|
||||
|
||||
## Goal Oriented
|
||||
|
||||
Steward does not work from prompts.
|
||||
|
||||
It works from goals.
|
||||
|
||||
Examples:
|
||||
|
||||
Maintain homelab reliability.
|
||||
|
||||
Prepare for house move.
|
||||
|
||||
Visit Queensland destinations before moving.
|
||||
|
||||
Reduce operational effort.
|
||||
|
||||
Research new technologies.
|
||||
|
||||
Complete backlog projects.
|
||||
|
||||
Prompts simply create, modify or clarify goals.
|
||||
|
||||
---
|
||||
|
||||
## Policy First
|
||||
|
||||
The AI never directly executes arbitrary commands.
|
||||
|
||||
Every action must pass through a policy engine.
|
||||
|
||||
Example policies:
|
||||
|
||||
Allowed automatically
|
||||
|
||||
- restart containers
|
||||
- retry backups
|
||||
- renew certificates
|
||||
- apply security patches
|
||||
|
||||
Requires approval
|
||||
|
||||
- Kubernetes upgrades
|
||||
- infrastructure changes
|
||||
- firewall modifications
|
||||
- snapshot deletion
|
||||
|
||||
Never
|
||||
|
||||
- delete backups
|
||||
- modify Steward permissions
|
||||
- disable audit logging
|
||||
|
||||
---
|
||||
|
||||
## Transparency
|
||||
|
||||
Every recommendation should explain:
|
||||
|
||||
- why
|
||||
- confidence
|
||||
- evidence
|
||||
- expected outcome
|
||||
- rollback strategy
|
||||
|
||||
No hidden reasoning.
|
||||
|
||||
---
|
||||
|
||||
## Conservative
|
||||
|
||||
Steward prefers:
|
||||
|
||||
- reversible actions
|
||||
- small changes
|
||||
- incremental improvement
|
||||
|
||||
rather than aggressive optimisation.
|
||||
|
||||
---
|
||||
|
||||
# Domains
|
||||
|
||||
Steward is intentionally domain-agnostic.
|
||||
|
||||
Initially it will focus on Homelab Operations.
|
||||
|
||||
Future domains include:
|
||||
|
||||
- family planning
|
||||
- travel
|
||||
- finance
|
||||
- home maintenance
|
||||
- learning
|
||||
- research
|
||||
- project management
|
||||
|
||||
Each domain has its own tools and permissions.
|
||||
|
||||
---
|
||||
|
||||
# Cognitive Load Reduction
|
||||
|
||||
The purpose is NOT task management.
|
||||
|
||||
The purpose is reducing the need to remember.
|
||||
|
||||
Traditional tools remember things once entered.
|
||||
|
||||
Steward helps remember to remember.
|
||||
|
||||
Examples
|
||||
|
||||
"I noticed you mentioned replacing the UPS several times."
|
||||
|
||||
"I noticed your dental check-up is overdue."
|
||||
|
||||
"I've researched family camping locations."
|
||||
|
||||
"I've identified a free weekend."
|
||||
|
||||
---
|
||||
|
||||
# Communication Philosophy
|
||||
|
||||
Steward communicates only when useful.
|
||||
|
||||
Communication categories:
|
||||
|
||||
Critical
|
||||
|
||||
Immediate interruption.
|
||||
|
||||
Decision Required
|
||||
|
||||
Requests human approval.
|
||||
|
||||
Background
|
||||
|
||||
No notification required.
|
||||
|
||||
Daily reports may exist but are not central.
|
||||
|
||||
Real-time context-aware communication is the default.
|
||||
|
||||
---
|
||||
|
||||
# Curiosity Budget
|
||||
|
||||
One of Steward's defining concepts.
|
||||
|
||||
If no operational work exists, Steward may spend limited compute researching or improving delegated objectives.
|
||||
|
||||
Examples
|
||||
|
||||
- Evaluate ArgoCD
|
||||
- Benchmark local LLMs
|
||||
- Research family holidays
|
||||
- Compare UPS replacements
|
||||
- Prototype monitoring improvements
|
||||
|
||||
Curiosity is budgeted.
|
||||
|
||||
The budget is a concrete, metered resource, not an aspiration. It is expressed in measurable units (for example: tokens/day, dollars/month, GPU-hours/week), set by the operator, and enforced by Steward refusing to exceed it. Consumption is observable so the operator can see what curiosity cost and what it produced.
|
||||
|
||||
It is suspended immediately if operational work becomes necessary.
|
||||
|
||||
---
|
||||
|
||||
# Confidence
|
||||
|
||||
Steward maintains historical confidence scores.
|
||||
|
||||
Example
|
||||
|
||||
Restart Jellyfin
|
||||
|
||||
Attempts: 34
|
||||
|
||||
Success: 33
|
||||
|
||||
Confidence: 97%
|
||||
|
||||
Upgrade Kubernetes
|
||||
|
||||
Attempts: 2
|
||||
|
||||
Success: 1
|
||||
|
||||
Rollback: 1
|
||||
|
||||
Confidence: 50%
|
||||
|
||||
Confidence is earned.
|
||||
|
||||
## Confidence and Autonomy
|
||||
|
||||
Confidence influences autonomy, but must never override policy.
|
||||
|
||||
Rules:
|
||||
|
||||
- The policy tier of an action (automatic / approval / never) is human-set and immutable to Steward. High confidence never promotes an action into a more permissive tier. Doing so would be self-modification by the back door.
|
||||
- Within a tier, confidence affects only *how* Steward acts: how strongly it recommends, how much evidence it attaches, whether it batches or surfaces individually.
|
||||
- Confidence is scoped, not global. "Restart Jellyfin" confidence does not transfer to "restart Postgres". Scope is (action × target × context), and Steward must not generalise across scopes without evidence.
|
||||
- Confidence decays. Environment changes (version upgrades, config changes, topology changes) invalidate historical success and should reset or discount the relevant scores. Stale confidence is treated as low confidence.
|
||||
|
||||
---
|
||||
|
||||
# Security Model
|
||||
|
||||
Safety is a property of the architecture, not just the model's good behaviour.
|
||||
|
||||
## The policy engine is the trust boundary
|
||||
|
||||
The policy engine is Steward's crown jewel and must be a separate, independently-audited component that the AI *calls*, not code the AI can read, modify, or bypass. The AI proposes actions; the policy engine decides. If the AI could edit the engine, every other guarantee collapses.
|
||||
|
||||
## Event content is untrusted
|
||||
|
||||
Events carry attacker-influenceable content: a GitHub PR title, a Telegram message, a webhook payload. Steward must treat all event content as untrusted input.
|
||||
|
||||
- Event content can inform situations and goals but can never escalate authority or promote an action's policy tier.
|
||||
- Instructions embedded in event content are data, not commands (defence against prompt injection).
|
||||
|
||||
## Memory can rot
|
||||
|
||||
"Nothing is forgotten" is both a feature and a liability. A learned preference can encode a mistake permanently, and learned lessons are themselves derived from potentially-untrusted history.
|
||||
|
||||
- Learned memory is reviewable and retractable. A bad lesson must be able to be inspected and removed.
|
||||
- New lessons that would change behaviour surface as *proposals*, consistent with [Reflection](#reflection), rather than silently altering future decisions.
|
||||
|
||||
---
|
||||
|
||||
# Memory
|
||||
|
||||
Different information belongs in different storage. The roles below are fixed; the named tools are current implementation choices, not commitments.
|
||||
|
||||
Structured store (e.g. NocoDB)
|
||||
|
||||
Examples
|
||||
|
||||
- Tasks
|
||||
- Goals
|
||||
- Incidents
|
||||
- Systems
|
||||
- Services
|
||||
- Confidence
|
||||
- Policies
|
||||
|
||||
Long-form store (e.g. vector database / RAG)
|
||||
|
||||
Examples
|
||||
|
||||
- documentation
|
||||
- conversations
|
||||
- runbooks
|
||||
- release notes
|
||||
|
||||
Metrics store (e.g. Prometheus)
|
||||
|
||||
Examples
|
||||
|
||||
- resource usage
|
||||
- trends
|
||||
- operational history
|
||||
|
||||
Audit
|
||||
|
||||
Immutable execution history.
|
||||
|
||||
Nothing is forgotten.
|
||||
|
||||
---
|
||||
|
||||
# Self Improvement
|
||||
|
||||
Steward should become better over time.
|
||||
|
||||
Important distinction
|
||||
|
||||
Self-improving
|
||||
|
||||
✅ Learn preferences
|
||||
|
||||
✅ Learn successful remediations
|
||||
|
||||
✅ Improve prompts
|
||||
|
||||
✅ Improve playbooks
|
||||
|
||||
✅ Improve planning
|
||||
|
||||
Self-modifying
|
||||
|
||||
❌ Rewrite policy engine
|
||||
|
||||
❌ Change permissions
|
||||
|
||||
❌ Remove safety controls
|
||||
|
||||
Immutable system components remain protected.
|
||||
|
||||
---
|
||||
|
||||
# Learning
|
||||
|
||||
Steward continuously learns:
|
||||
|
||||
Operator preferences
|
||||
|
||||
Family preferences
|
||||
|
||||
Infrastructure behaviour
|
||||
|
||||
Recurring incidents
|
||||
|
||||
Successful remediation
|
||||
|
||||
Failed remediation
|
||||
|
||||
Communication preferences
|
||||
|
||||
Historical context
|
||||
|
||||
This learning improves planning.
|
||||
|
||||
---
|
||||
|
||||
# Reflection
|
||||
|
||||
Periodic reflection sessions.
|
||||
|
||||
Questions include:
|
||||
|
||||
What repeated?
|
||||
|
||||
What wasted time?
|
||||
|
||||
What should be automated?
|
||||
|
||||
Which playbooks failed?
|
||||
|
||||
Which policies need review?
|
||||
|
||||
Which research produced value?
|
||||
|
||||
Reflection generates proposals rather than automatic change.
|
||||
|
||||
---
|
||||
|
||||
# Research Workflow
|
||||
|
||||
Research tasks behave like engineering work.
|
||||
|
||||
Example
|
||||
|
||||
Research destination --> Gather information --> Summarise findings --> Generate recommendations --> Present proposal --> Receive feedback --> Update understanding --> Re-plan
|
||||
|
||||
---
|
||||
|
||||
# Engineering Workflow
|
||||
|
||||
Operational improvements follow software engineering practices.
|
||||
|
||||
Observe issue --> Generate proposal --> Approval --> Create Git branch --> Modify playbooks --> Run tests --> Open Pull Request --> Deploy through AWX --> Observe --> Learn
|
||||
|
||||
---
|
||||
|
||||
# Long-term Vision
|
||||
|
||||
Steward evolves from:
|
||||
|
||||
Assistant --> Operator --> Engineer --> Trusted Steward
|
||||
|
||||
Trust is earned through demonstrated reliability rather than increasing model capability.
|
||||
|
||||
---
|
||||
|
||||
# Success Criteria
|
||||
|
||||
Steward is successful when:
|
||||
|
||||
I think less.
|
||||
|
||||
I forget less.
|
||||
|
||||
Routine work disappears.
|
||||
|
||||
The homelab becomes increasingly self-maintaining.
|
||||
|
||||
Projects continue progressing while I am busy.
|
||||
|
||||
The system accumulates operational knowledge.
|
||||
|
||||
The system proactively suggests valuable improvements.
|
||||
|
||||
The system remains predictable.
|
||||
|
||||
The system remains explainable.
|
||||
|
||||
The system remains safe.
|
||||
|
||||
## Measurable Proxies
|
||||
|
||||
The criteria above are subjective. To tell whether v0.2 actually beat v0.1, track measurable proxies alongside them:
|
||||
|
||||
- Operator interruptions per week (trend down).
|
||||
- Ratio of actions auto-approved vs escalated for approval.
|
||||
- Mean time to remediation for recurring incidents.
|
||||
- False-alarm / unnecessary-notification rate.
|
||||
- Curiosity spend vs value produced (proposals accepted).
|
||||
|
||||
---
|
||||
|
||||
# Mission Statement
|
||||
|
||||
Steward exists to maximise the reliability, maintainability and usefulness of its delegated domains while minimising operator cognitive load.
|
||||
|
||||
It continuously observes, remembers, researches, plans and acts within explicitly delegated authority.
|
||||
|
||||
It values safety over speed, explanation over opacity, and long-term trust over short-term autonomy.
|
||||
14
docker-compose.yml
Normal file
14
docker-compose.yml
Normal file
@ -0,0 +1,14 @@
|
||||
services:
|
||||
steward:
|
||||
image: ghcr.io/djw4/steward:latest
|
||||
# To build locally instead: uncomment the next line and comment out image above
|
||||
# build: .
|
||||
user: "1000:1000" # matches the UID/GID created in the Dockerfile
|
||||
env_file: .env # copy .env.example → .env and fill in your values
|
||||
environment:
|
||||
THREAD_MEMORY_PATH: /data/thread_memory.json
|
||||
volumes:
|
||||
# Bind-mount a local ./data directory for persistent storage.
|
||||
# Create it before the first run: mkdir -p data
|
||||
- ./data:/data
|
||||
restart: unless-stopped
|
||||
52
pyproject.toml
Normal file
52
pyproject.toml
Normal file
@ -0,0 +1,52 @@
|
||||
[build-system]
|
||||
requires = ["setuptools>=68", "wheel"]
|
||||
build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "steward"
|
||||
version = "0.1.0"
|
||||
description = "A long-running, AI-assisted personal operations platform"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.12"
|
||||
dependencies = [
|
||||
"python-telegram-bot>=21.0",
|
||||
"openai>=1.30",
|
||||
"httpx>=0.27",
|
||||
"python-dotenv>=1.0",
|
||||
"apscheduler>=3.10",
|
||||
"pydantic>=2.7",
|
||||
"pydantic-settings>=2.3",
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
dev = [
|
||||
"pytest>=8.2",
|
||||
"pytest-asyncio>=0.23",
|
||||
"pytest-mock>=3.14",
|
||||
"ruff>=0.4",
|
||||
"mypy>=1.10",
|
||||
"respx>=0.21",
|
||||
]
|
||||
|
||||
[project.scripts]
|
||||
steward = "steward.main:main"
|
||||
|
||||
[tool.setuptools.packages.find]
|
||||
where = ["."]
|
||||
include = ["steward*"]
|
||||
|
||||
[tool.ruff]
|
||||
line-length = 100
|
||||
target-version = "py312"
|
||||
|
||||
[tool.ruff.lint]
|
||||
select = ["E", "F", "I", "UP"]
|
||||
|
||||
[tool.mypy]
|
||||
python_version = "3.12"
|
||||
strict = true
|
||||
ignore_missing_imports = true
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
asyncio_mode = "auto"
|
||||
testpaths = ["tests"]
|
||||
1
steward/__init__.py
Normal file
1
steward/__init__.py
Normal file
@ -0,0 +1 @@
|
||||
"""Steward – AI-assisted personal operations platform."""
|
||||
0
steward/bot/__init__.py
Normal file
0
steward/bot/__init__.py
Normal file
428
steward/bot/telegram.py
Normal file
428
steward/bot/telegram.py
Normal file
@ -0,0 +1,428 @@
|
||||
"""Telegram bot interface for Steward."""
|
||||
|
||||
import logging
|
||||
from collections import defaultdict
|
||||
from typing import Any
|
||||
|
||||
from telegram import Update
|
||||
from telegram.constants import ParseMode
|
||||
from telegram.ext import (
|
||||
Application,
|
||||
CommandHandler,
|
||||
ContextTypes,
|
||||
MessageHandler,
|
||||
filters,
|
||||
)
|
||||
|
||||
from steward.config import Settings
|
||||
from steward.llm.client import LLMClient
|
||||
from steward.memory.thread_store import ThreadMemoryStore, ThreadSummary
|
||||
from steward.proposals.generator import Proposal, ProposalGenerator
|
||||
from steward.tools.client import ToolClient
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Per-user conversation history for non-threaded messages (capped at _MAX_HISTORY turns).
|
||||
_history: dict[int, list[dict[str, Any]]] = defaultdict(list)
|
||||
_MAX_HISTORY = 20
|
||||
|
||||
# Per-thread conversation history: (chat_id, thread_id) → full history (unbounded).
|
||||
# Messages belonging to a Telegram message thread are kept in their entirety here
|
||||
# until explicitly flushed by the /flush command.
|
||||
_thread_history: dict[tuple[int, int], list[dict[str, Any]]] = defaultdict(list)
|
||||
|
||||
_FLUSH_SYSTEM_PROMPT = (
|
||||
"You are Steward. The following is a complete Telegram message thread conversation. "
|
||||
"Produce a concise but comprehensive summary that captures:\n"
|
||||
"- The main topics discussed\n"
|
||||
"- Key decisions or conclusions reached\n"
|
||||
"- Any outstanding actions or open questions\n"
|
||||
"- Important context that would help recall this conversation later\n\n"
|
||||
"Be precise. Omit pleasantries."
|
||||
)
|
||||
|
||||
_TAGS_SYSTEM_PROMPT = (
|
||||
"You are a keyword tagger for a knowledge base. "
|
||||
"Extract 5–8 short, lowercase keyword tags from the following conversation summary. "
|
||||
"Tags should represent the main topics, entities, and concepts discussed. "
|
||||
"Return ONLY a comma-separated list of tags with no other text or punctuation. "
|
||||
"Example output: api design, authentication, database schema, user roles, caching"
|
||||
)
|
||||
|
||||
_KB_CONTEXT_HEADER = (
|
||||
"The following are relevant past conversation summaries from your knowledge base. "
|
||||
"Use them as background context if they relate to the current question, "
|
||||
"but do not repeat their contents unless directly asked."
|
||||
)
|
||||
|
||||
|
||||
def _thread_key(update: Update) -> tuple[int, int] | None:
|
||||
"""Return the (chat_id, thread_id) key if the message is part of a thread, else None."""
|
||||
msg = update.message
|
||||
chat = update.effective_chat
|
||||
if msg is None or chat is None:
|
||||
return None
|
||||
thread_id = msg.message_thread_id
|
||||
if thread_id is None:
|
||||
return None
|
||||
return (chat.id, thread_id)
|
||||
|
||||
|
||||
def _is_allowed(user_id: int, settings: Settings) -> bool:
|
||||
"""Return True if the user is in the allow-list (or no list is configured)."""
|
||||
if not settings.telegram_allowed_user_ids:
|
||||
return True
|
||||
return user_id in settings.telegram_allowed_user_ids
|
||||
|
||||
|
||||
async def _send_long(update: Update, text: str) -> None:
|
||||
"""Send text, splitting across messages if it exceeds Telegram's 4096-char limit."""
|
||||
limit = 4096
|
||||
for i in range(0, len(text), limit):
|
||||
await update.message.reply_text( # type: ignore[union-attr]
|
||||
text[i : i + limit],
|
||||
parse_mode=ParseMode.MARKDOWN,
|
||||
)
|
||||
|
||||
|
||||
async def start_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
|
||||
"""Handle /start."""
|
||||
settings: Settings = context.bot_data["settings"]
|
||||
user = update.effective_user
|
||||
if user is None or not _is_allowed(user.id, settings):
|
||||
return
|
||||
|
||||
await update.message.reply_text( # type: ignore[union-attr]
|
||||
"Hello, I'm *Steward* \U0001f916\n\n"
|
||||
"I'm your AI-assisted personal operations platform.\n"
|
||||
"Talk to me naturally, or use:\n"
|
||||
"/help – show available commands\n"
|
||||
"/clear – reset conversation history\n"
|
||||
"/flush – summarise and archive this thread's memory\n"
|
||||
"/recall [query] – retrieve archived thread summaries\n"
|
||||
"/analyse – run a manual API analysis right now",
|
||||
parse_mode=ParseMode.MARKDOWN,
|
||||
)
|
||||
|
||||
|
||||
async def help_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
|
||||
"""Handle /help."""
|
||||
settings: Settings = context.bot_data["settings"]
|
||||
user = update.effective_user
|
||||
if user is None or not _is_allowed(user.id, settings):
|
||||
return
|
||||
|
||||
await update.message.reply_text( # type: ignore[union-attr]
|
||||
"*Steward commands*\n\n"
|
||||
"/start – greeting\n"
|
||||
"/help – this message\n"
|
||||
"/clear – reset conversation history for this context\n"
|
||||
"/flush – summarise the current thread, store the summary, and compress memory\n"
|
||||
" _(only available inside a message thread)_\n"
|
||||
"/recall [query] – show this thread's summary, list all summaries, or search by keyword\n"
|
||||
"/analyse – trigger an immediate API analysis and proposal",
|
||||
parse_mode=ParseMode.MARKDOWN,
|
||||
)
|
||||
|
||||
|
||||
async def clear_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
|
||||
"""Handle /clear – wipe conversation history for this context.
|
||||
|
||||
Inside a message thread: clears the thread's unbounded history.
|
||||
Outside a thread: clears the per-user capped history.
|
||||
"""
|
||||
settings: Settings = context.bot_data["settings"]
|
||||
user = update.effective_user
|
||||
if user is None or not _is_allowed(user.id, settings):
|
||||
return
|
||||
|
||||
key = _thread_key(update)
|
||||
if key is not None:
|
||||
_thread_history[key].clear()
|
||||
await update.message.reply_text( # type: ignore[union-attr]
|
||||
"Thread conversation history cleared."
|
||||
)
|
||||
else:
|
||||
_history[user.id].clear()
|
||||
await update.message.reply_text( # type: ignore[union-attr]
|
||||
"Conversation history cleared."
|
||||
)
|
||||
|
||||
|
||||
async def flush_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
|
||||
"""Handle /flush – summarise thread memory, persist it, and compress in-memory history.
|
||||
|
||||
Steps:
|
||||
1. Verify the command is issued inside a message thread.
|
||||
2. Summarise the full thread history via the LLM.
|
||||
3. Persist the summary in the ThreadMemoryStore (keyed by chat_id + thread_id).
|
||||
4. Replace the in-memory thread history with a single compressed context message
|
||||
so conversation can continue with the summary as background.
|
||||
"""
|
||||
settings: Settings = context.bot_data["settings"]
|
||||
llm: LLMClient = context.bot_data["llm"]
|
||||
store: ThreadMemoryStore = context.bot_data["thread_store"]
|
||||
user = update.effective_user
|
||||
|
||||
if user is None or not _is_allowed(user.id, settings):
|
||||
return
|
||||
|
||||
key = _thread_key(update)
|
||||
if key is None:
|
||||
await update.message.reply_text( # type: ignore[union-attr]
|
||||
"\u26a0\ufe0f /flush can only be used inside a message thread."
|
||||
)
|
||||
return
|
||||
|
||||
chat_id, thread_id = key
|
||||
history = _thread_history[key]
|
||||
|
||||
if not history:
|
||||
await update.message.reply_text( # type: ignore[union-attr]
|
||||
"This thread has no conversation history to flush."
|
||||
)
|
||||
return
|
||||
|
||||
await update.message.reply_text( # type: ignore[union-attr]
|
||||
"\U0001f4be Summarising thread memory\u2026"
|
||||
)
|
||||
|
||||
# Build a readable transcript for the LLM to summarise
|
||||
transcript_lines = []
|
||||
for msg in history:
|
||||
role_label = "User" if msg["role"] == "user" else "Steward"
|
||||
transcript_lines.append(f"{role_label}: {msg['content']}")
|
||||
transcript = "\n".join(transcript_lines)
|
||||
|
||||
summary_text = await llm.chat(
|
||||
f"Thread transcript:\n\n{transcript}",
|
||||
system_prompt=_FLUSH_SYSTEM_PROMPT,
|
||||
)
|
||||
|
||||
# Extract keyword tags for knowledge-base indexing (second LLM call, lightweight)
|
||||
tags_raw = await llm.chat(
|
||||
f"Summary to tag:\n\n{summary_text}",
|
||||
system_prompt=_TAGS_SYSTEM_PROMPT,
|
||||
)
|
||||
tags = [t.strip().lower() for t in tags_raw.split(",") if t.strip()][:8]
|
||||
|
||||
message_count = sum(1 for m in history if m["role"] == "user")
|
||||
thread_summary = ThreadSummary(
|
||||
chat_id=chat_id,
|
||||
thread_id=thread_id,
|
||||
summary=summary_text,
|
||||
message_count=message_count,
|
||||
tags=tags,
|
||||
)
|
||||
store.save(thread_summary)
|
||||
|
||||
# Compress: replace history with a single system-context entry so the thread
|
||||
# can continue with the summary as background knowledge.
|
||||
_thread_history[key] = [
|
||||
{"role": "system", "content": f"Summary of earlier conversation:\n{summary_text}"}
|
||||
]
|
||||
|
||||
await update.message.reply_text( # type: ignore[union-attr]
|
||||
f"\u2705 Thread memory flushed and stored "
|
||||
f"(thread `{thread_id}`, {message_count} messages summarised).",
|
||||
parse_mode=ParseMode.MARKDOWN,
|
||||
)
|
||||
|
||||
|
||||
async def recall_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
|
||||
"""Handle /recall [query] – retrieve stored thread summaries.
|
||||
|
||||
With a query argument (e.g. ``/recall api design``): searches all stored
|
||||
summaries whose tags overlap with the query keywords and returns matches.
|
||||
|
||||
Without a query:
|
||||
- Inside a thread: shows the stored summary for this thread (if any).
|
||||
- Outside a thread: lists all stored summaries (newest first).
|
||||
"""
|
||||
settings: Settings = context.bot_data["settings"]
|
||||
store: ThreadMemoryStore = context.bot_data["thread_store"]
|
||||
user = update.effective_user
|
||||
|
||||
if user is None or not _is_allowed(user.id, settings):
|
||||
return
|
||||
|
||||
# If the user supplied a keyword query, search the knowledge base
|
||||
args: list[str] = context.args or [] # type: ignore[assignment]
|
||||
if args:
|
||||
query = " ".join(args).strip()
|
||||
results = store.search(query)
|
||||
if not results:
|
||||
await update.message.reply_text( # type: ignore[union-attr]
|
||||
f"No memories found matching *{query}*. "
|
||||
"Try a different keyword or use /flush to add more summaries.",
|
||||
parse_mode=ParseMode.MARKDOWN,
|
||||
)
|
||||
return
|
||||
lines = [f"\U0001f50d *Knowledge base search: {query}*\n"]
|
||||
for s in results:
|
||||
date = s.flushed_at[:10]
|
||||
first_line = s.summary.split("\n")[0][:80]
|
||||
lines.append(f"\u2022 Thread `{s.thread_id}` ({date}): {first_line}\u2026")
|
||||
if s.tags:
|
||||
lines.append(f" \U0001f3f7 {', '.join(s.tags)}")
|
||||
await _send_long(update, "\n".join(lines))
|
||||
return
|
||||
|
||||
key = _thread_key(update)
|
||||
if key is not None:
|
||||
chat_id, thread_id = key
|
||||
stored = store.get(chat_id, thread_id)
|
||||
if stored is None:
|
||||
await update.message.reply_text( # type: ignore[union-attr]
|
||||
"No stored summary for this thread yet. Use /flush to create one."
|
||||
)
|
||||
else:
|
||||
await _send_long(update, stored.format_for_telegram())
|
||||
return
|
||||
|
||||
# Outside a thread: list all stored summaries
|
||||
all_summaries = store.all()
|
||||
if not all_summaries:
|
||||
await update.message.reply_text( # type: ignore[union-attr]
|
||||
"No thread summaries stored yet. Use /flush inside a message thread."
|
||||
)
|
||||
return
|
||||
|
||||
lines = ["\U0001f4da *Stored thread summaries*\n"]
|
||||
for s in all_summaries:
|
||||
date = s.flushed_at[:10]
|
||||
first_line = s.summary.split("\n")[0][:80]
|
||||
lines.append(f"\u2022 Thread `{s.thread_id}` ({date}): {first_line}\u2026")
|
||||
if s.tags:
|
||||
lines.append(f" \U0001f3f7 {', '.join(s.tags)}")
|
||||
await _send_long(update, "\n".join(lines))
|
||||
|
||||
|
||||
async def analyse_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
|
||||
"""Handle /analyse – run the proposal generator on demand."""
|
||||
settings: Settings = context.bot_data["settings"]
|
||||
llm: LLMClient = context.bot_data["llm"]
|
||||
user = update.effective_user
|
||||
if user is None or not _is_allowed(user.id, settings):
|
||||
return
|
||||
|
||||
await update.message.reply_text("Running analysis, please wait...") # type: ignore[union-attr]
|
||||
generator = ProposalGenerator(settings, llm)
|
||||
proposal = await generator.run()
|
||||
if proposal is None:
|
||||
await update.message.reply_text( # type: ignore[union-attr]
|
||||
"Analysis could not be completed. "
|
||||
"Check that `ANALYSIS_TARGET_URL` is configured."
|
||||
)
|
||||
return
|
||||
|
||||
await _send_long(update, proposal.format_for_telegram())
|
||||
|
||||
|
||||
def _with_kb_context(
|
||||
history: list[dict[str, Any]],
|
||||
relevant: list[ThreadSummary],
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Prepend relevant knowledge-base summaries as a transient system context message.
|
||||
|
||||
The returned list is a *new* list — the original *history* is not mutated.
|
||||
The injected message is never appended to the stored history, so it does not
|
||||
permanently consume the context window.
|
||||
"""
|
||||
if not relevant:
|
||||
return history
|
||||
snippets = [f"[Thread {s.thread_id}] {s.summary[:400]}" for s in relevant[:3]]
|
||||
kb_msg = _KB_CONTEXT_HEADER + "\n\n" + "\n\n---\n\n".join(snippets)
|
||||
return [{"role": "system", "content": kb_msg}, *history]
|
||||
|
||||
|
||||
async def message_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
|
||||
"""Handle plain text messages – forward to LLM and reply.
|
||||
|
||||
Thread messages: history is stored unbounded under the (chat_id, thread_id) key.
|
||||
Non-thread messages: history is capped at _MAX_HISTORY turns per user.
|
||||
|
||||
Before each LLM call the knowledge base is searched for summaries whose tags
|
||||
overlap with keywords in the current message. Any matches are injected as
|
||||
transient context — they are NOT stored in the rolling history, so they do
|
||||
not permanently consume the context window.
|
||||
|
||||
If a :class:`~steward.tools.client.ToolClient` is registered in
|
||||
``context.bot_data``, the LLM is invoked with tool calling support so it
|
||||
can take actions on the configured MCP/OpenAPI tool server.
|
||||
"""
|
||||
settings: Settings = context.bot_data["settings"]
|
||||
llm: LLMClient = context.bot_data["llm"]
|
||||
store: ThreadMemoryStore = context.bot_data["thread_store"]
|
||||
tool_client: ToolClient | None = context.bot_data.get("tool_client")
|
||||
user = update.effective_user
|
||||
if user is None or not _is_allowed(user.id, settings):
|
||||
return
|
||||
|
||||
text = update.message.text # type: ignore[union-attr]
|
||||
if not text:
|
||||
return
|
||||
|
||||
key = _thread_key(update)
|
||||
if key is not None:
|
||||
# Thread message: unbounded history
|
||||
history: list[dict[str, Any]] = _thread_history[key]
|
||||
call_history = _with_kb_context(history, store.search(text))
|
||||
if tool_client is not None:
|
||||
reply = await llm.chat_with_tools(text, tool_client, history=call_history)
|
||||
else:
|
||||
reply = await llm.chat(text, history=call_history)
|
||||
history.append({"role": "user", "content": text})
|
||||
history.append({"role": "assistant", "content": reply})
|
||||
else:
|
||||
# Non-thread message: capped history per user
|
||||
user_history: list[dict[str, Any]] = _history[user.id]
|
||||
call_history = _with_kb_context(user_history, store.search(text))
|
||||
if tool_client is not None:
|
||||
reply = await llm.chat_with_tools(text, tool_client, history=call_history)
|
||||
else:
|
||||
reply = await llm.chat(text, history=call_history)
|
||||
user_history.append({"role": "user", "content": text})
|
||||
user_history.append({"role": "assistant", "content": reply})
|
||||
if len(user_history) > _MAX_HISTORY * 2:
|
||||
_history[user.id] = user_history[-(_MAX_HISTORY * 2) :]
|
||||
|
||||
await _send_long(update, reply)
|
||||
|
||||
|
||||
def build_application(
|
||||
settings: Settings,
|
||||
llm: LLMClient,
|
||||
thread_store: ThreadMemoryStore | None = None,
|
||||
tool_client: ToolClient | None = None,
|
||||
) -> Application: # type: ignore[type-arg]
|
||||
"""Build and return the Telegram Application."""
|
||||
app = Application.builder().token(settings.telegram_bot_token).build()
|
||||
app.bot_data["settings"] = settings
|
||||
app.bot_data["llm"] = llm
|
||||
app.bot_data["thread_store"] = thread_store or ThreadMemoryStore(settings.thread_memory_path)
|
||||
app.bot_data["tool_client"] = tool_client # None when tools are not configured
|
||||
|
||||
app.add_handler(CommandHandler("start", start_handler))
|
||||
app.add_handler(CommandHandler("help", help_handler))
|
||||
app.add_handler(CommandHandler("clear", clear_handler))
|
||||
app.add_handler(CommandHandler("flush", flush_handler))
|
||||
app.add_handler(CommandHandler("recall", recall_handler))
|
||||
app.add_handler(CommandHandler("analyse", analyse_handler))
|
||||
app.add_handler(MessageHandler(filters.TEXT & ~filters.COMMAND, message_handler))
|
||||
|
||||
return app
|
||||
|
||||
|
||||
async def send_proposal(app: Application, proposal: Proposal, user_ids: list[int]) -> None: # type: ignore[type-arg]
|
||||
"""Send a proposal message to all configured user IDs."""
|
||||
text = proposal.format_for_telegram()
|
||||
for uid in user_ids:
|
||||
try:
|
||||
await app.bot.send_message(
|
||||
chat_id=uid,
|
||||
text=text,
|
||||
parse_mode=ParseMode.MARKDOWN,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("Failed to send proposal to user %d", uid)
|
||||
44
steward/config.py
Normal file
44
steward/config.py
Normal file
@ -0,0 +1,44 @@
|
||||
"""Configuration management for Steward."""
|
||||
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
|
||||
|
||||
class Settings(BaseSettings):
|
||||
"""Application settings loaded from environment variables or .env file."""
|
||||
|
||||
model_config = SettingsConfigDict(env_file=".env", env_file_encoding="utf-8", extra="ignore")
|
||||
|
||||
# Telegram
|
||||
telegram_bot_token: str = ""
|
||||
telegram_allowed_user_ids: list[int] = []
|
||||
|
||||
# LLM (OpenAI-compatible)
|
||||
openai_api_key: str = ""
|
||||
openai_base_url: str = "https://api.openai.com/v1"
|
||||
openai_model: str = "gpt-4o"
|
||||
openai_system_prompt: str = (
|
||||
"You are Steward, a persistent, trustworthy AI-assisted personal operations platform. "
|
||||
"You reduce cognitive load by observing, remembering, planning, and proposing actions. "
|
||||
"You are conservative, transparent, and policy-aware. "
|
||||
"Always explain your reasoning."
|
||||
)
|
||||
|
||||
# Analysis / proposal worker
|
||||
analysis_target_url: str = ""
|
||||
analysis_target_api_key: str = ""
|
||||
analysis_cron_hour: int = 8
|
||||
analysis_cron_minute: int = 0
|
||||
|
||||
# Thread memory
|
||||
thread_memory_path: str = "thread_memory.json"
|
||||
|
||||
# MCP / OpenAPI tool server (open-webui/openapi-servers compatible)
|
||||
# Set MCP_SERVER_URL to enable tool calling. The service fetches
|
||||
# /openapi.json from this URL to discover available tools.
|
||||
mcp_server_url: str = ""
|
||||
mcp_server_api_key: str = ""
|
||||
|
||||
|
||||
def get_settings() -> Settings:
|
||||
"""Return application settings singleton."""
|
||||
return Settings()
|
||||
0
steward/llm/__init__.py
Normal file
0
steward/llm/__init__.py
Normal file
179
steward/llm/client.py
Normal file
179
steward/llm/client.py
Normal file
@ -0,0 +1,179 @@
|
||||
"""LLM client wrapper (OpenAI-compatible)."""
|
||||
|
||||
import json
|
||||
import logging
|
||||
from collections.abc import AsyncIterator
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from openai import AsyncOpenAI
|
||||
|
||||
from steward.config import Settings
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from steward.tools.client import ToolClient
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class LLMClient:
|
||||
"""Thin async wrapper around the OpenAI chat-completions API."""
|
||||
|
||||
def __init__(self, settings: Settings) -> None:
|
||||
self._settings = settings
|
||||
self._client = AsyncOpenAI(
|
||||
api_key=settings.openai_api_key,
|
||||
base_url=settings.openai_base_url,
|
||||
)
|
||||
|
||||
async def chat(
|
||||
self,
|
||||
user_message: str,
|
||||
*,
|
||||
history: list[dict[str, Any]] | None = None,
|
||||
system_prompt: str | None = None,
|
||||
) -> str:
|
||||
"""Send a user message (with optional history) and return the assistant reply."""
|
||||
system = system_prompt or self._settings.openai_system_prompt
|
||||
messages: list[dict[str, Any]] = [{"role": "system", "content": system}]
|
||||
if history:
|
||||
messages.extend(history)
|
||||
messages.append({"role": "user", "content": user_message})
|
||||
|
||||
logger.debug(
|
||||
"LLM request: model=%s messages=%d", self._settings.openai_model, len(messages)
|
||||
)
|
||||
response = await self._client.chat.completions.create(
|
||||
model=self._settings.openai_model,
|
||||
messages=messages, # type: ignore[arg-type]
|
||||
)
|
||||
reply = response.choices[0].message.content or ""
|
||||
logger.debug("LLM reply: %d chars", len(reply))
|
||||
return reply
|
||||
|
||||
async def chat_with_tools(
|
||||
self,
|
||||
user_message: str,
|
||||
tool_client: "ToolClient",
|
||||
*,
|
||||
history: list[dict[str, Any]] | None = None,
|
||||
system_prompt: str | None = None,
|
||||
max_tool_rounds: int = 10,
|
||||
) -> str:
|
||||
"""Chat with tool calling support (agentic loop).
|
||||
|
||||
Sends *user_message* to the LLM with the tool definitions supplied by
|
||||
*tool_client*. If the model requests one or more tool calls, each call
|
||||
is executed via *tool_client* and the results are fed back to the model.
|
||||
The loop repeats until the model returns a plain-text response or
|
||||
*max_tool_rounds* is reached.
|
||||
|
||||
Intermediate tool-call messages are **not** returned to the caller and
|
||||
should **not** be stored in the persistent conversation history; only
|
||||
the final assistant reply needs to be recorded alongside the original
|
||||
user message.
|
||||
"""
|
||||
tools = await tool_client.get_tools()
|
||||
system = system_prompt or self._settings.openai_system_prompt
|
||||
messages: list[dict[str, Any]] = [{"role": "system", "content": system}]
|
||||
if history:
|
||||
messages.extend(history)
|
||||
messages.append({"role": "user", "content": user_message})
|
||||
|
||||
logger.debug(
|
||||
"LLM tool-call request: model=%s tools=%d messages=%d",
|
||||
self._settings.openai_model,
|
||||
len(tools),
|
||||
len(messages),
|
||||
)
|
||||
|
||||
last_content = ""
|
||||
for round_num in range(max_tool_rounds):
|
||||
kwargs: dict[str, Any] = {
|
||||
"model": self._settings.openai_model,
|
||||
"messages": messages,
|
||||
}
|
||||
if tools:
|
||||
kwargs["tools"] = tools
|
||||
kwargs["tool_choice"] = "auto"
|
||||
|
||||
response = await self._client.chat.completions.create(**kwargs) # type: ignore[arg-type]
|
||||
choice = response.choices[0]
|
||||
msg = choice.message
|
||||
last_content = msg.content or ""
|
||||
|
||||
if not msg.tool_calls:
|
||||
# No tool calls → final answer
|
||||
logger.debug("LLM reply (round %d): %d chars", round_num, len(last_content))
|
||||
return last_content
|
||||
|
||||
tool_names = [tc.function.name for tc in msg.tool_calls]
|
||||
logger.info("Tool calls in round %d: %s", round_num + 1, tool_names)
|
||||
|
||||
# Append the assistant's tool-call turn
|
||||
messages.append(
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": msg.content,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": tc.id,
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": tc.function.name,
|
||||
"arguments": tc.function.arguments,
|
||||
},
|
||||
}
|
||||
for tc in msg.tool_calls
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
# Execute each tool and append results
|
||||
for tc in msg.tool_calls:
|
||||
try:
|
||||
try:
|
||||
args = json.loads(tc.function.arguments)
|
||||
except json.JSONDecodeError as exc:
|
||||
raise ValueError(
|
||||
f"Invalid JSON in tool arguments for '{tc.function.name}': {exc}"
|
||||
) from exc
|
||||
result = await tool_client.call(tc.function.name, args)
|
||||
logger.info("Tool %s → %d chars", tc.function.name, len(result))
|
||||
except Exception as exc:
|
||||
result = f"Tool call failed: {exc}"
|
||||
logger.warning("Tool %s failed: %s", tc.function.name, exc)
|
||||
|
||||
messages.append(
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": tc.id,
|
||||
"content": result,
|
||||
}
|
||||
)
|
||||
|
||||
logger.warning("Reached max tool rounds (%d); returning last content", max_tool_rounds)
|
||||
return last_content
|
||||
|
||||
async def stream(
|
||||
self,
|
||||
user_message: str,
|
||||
*,
|
||||
history: list[dict[str, Any]] | None = None,
|
||||
system_prompt: str | None = None,
|
||||
) -> AsyncIterator[str]:
|
||||
"""Stream the assistant reply token by token."""
|
||||
system = system_prompt or self._settings.openai_system_prompt
|
||||
messages: list[dict[str, Any]] = [{"role": "system", "content": system}]
|
||||
if history:
|
||||
messages.extend(history)
|
||||
messages.append({"role": "user", "content": user_message})
|
||||
|
||||
stream = await self._client.chat.completions.create(
|
||||
model=self._settings.openai_model,
|
||||
messages=messages, # type: ignore[arg-type]
|
||||
stream=True,
|
||||
)
|
||||
async for chunk in stream:
|
||||
delta = chunk.choices[0].delta.content
|
||||
if delta:
|
||||
yield delta
|
||||
84
steward/main.py
Normal file
84
steward/main.py
Normal file
@ -0,0 +1,84 @@
|
||||
"""Steward application entry point."""
|
||||
|
||||
import logging
|
||||
import sys
|
||||
|
||||
from apscheduler.schedulers.asyncio import AsyncIOScheduler
|
||||
|
||||
from steward.bot.telegram import build_application, send_proposal
|
||||
from steward.config import get_settings
|
||||
from steward.llm.client import LLMClient
|
||||
from steward.memory.thread_store import ThreadMemoryStore
|
||||
from steward.proposals.generator import ProposalGenerator
|
||||
from steward.tools.client import ToolClient
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format="%(asctime)s %(levelname)s %(name)s: %(message)s",
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def _run_scheduled_analysis(
|
||||
generator: ProposalGenerator,
|
||||
app, # type: ignore[type-arg]
|
||||
user_ids: list[int],
|
||||
) -> None:
|
||||
"""Scheduled job: run analysis and push proposal via Telegram."""
|
||||
logger.info("Running scheduled analysis")
|
||||
proposal = await generator.run()
|
||||
if proposal:
|
||||
await send_proposal(app, proposal, user_ids)
|
||||
else:
|
||||
logger.info("No proposal generated (generator returned None)")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
"""Start Steward."""
|
||||
settings = get_settings()
|
||||
|
||||
if not settings.telegram_bot_token:
|
||||
logger.error("TELEGRAM_BOT_TOKEN is not set – cannot start")
|
||||
sys.exit(1)
|
||||
|
||||
if not settings.openai_api_key:
|
||||
logger.error("OPENAI_API_KEY is not set – cannot start")
|
||||
sys.exit(1)
|
||||
|
||||
llm = LLMClient(settings)
|
||||
thread_store = ThreadMemoryStore(settings.thread_memory_path)
|
||||
|
||||
tool_client: ToolClient | None = None
|
||||
if settings.mcp_server_url:
|
||||
tool_client = ToolClient(
|
||||
base_url=settings.mcp_server_url,
|
||||
api_key=settings.mcp_server_api_key,
|
||||
)
|
||||
logger.info("MCP tool server configured: %s", settings.mcp_server_url)
|
||||
else:
|
||||
logger.info("No MCP_SERVER_URL configured – tool calling disabled")
|
||||
|
||||
app = build_application(settings, llm, thread_store, tool_client)
|
||||
generator = ProposalGenerator(settings, llm)
|
||||
|
||||
scheduler = AsyncIOScheduler()
|
||||
scheduler.add_job(
|
||||
_run_scheduled_analysis,
|
||||
"cron",
|
||||
hour=settings.analysis_cron_hour,
|
||||
minute=settings.analysis_cron_minute,
|
||||
args=[generator, app, settings.telegram_allowed_user_ids],
|
||||
)
|
||||
scheduler.start()
|
||||
|
||||
logger.info(
|
||||
"Steward starting: model=%s analysis_url=%s tools=%s",
|
||||
settings.openai_model,
|
||||
settings.analysis_target_url or "(none)",
|
||||
settings.mcp_server_url or "(none)",
|
||||
)
|
||||
app.run_polling(allowed_updates=["message"])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
0
steward/memory/__init__.py
Normal file
0
steward/memory/__init__.py
Normal file
121
steward/memory/thread_store.py
Normal file
121
steward/memory/thread_store.py
Normal file
@ -0,0 +1,121 @@
|
||||
"""Persistent storage for flushed thread summaries (knowledge-base style)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
from dataclasses import asdict, dataclass, field
|
||||
from datetime import UTC, datetime
|
||||
from pathlib import Path
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ThreadSummary:
|
||||
"""A persisted summary of a flushed Telegram message thread.
|
||||
|
||||
``tags`` is a list of short lowercase keywords extracted by the LLM at flush
|
||||
time. They are used to index the knowledge base so summaries can be recalled
|
||||
contextually without being kept permanently in the conversation context.
|
||||
"""
|
||||
|
||||
chat_id: int
|
||||
thread_id: int
|
||||
summary: str
|
||||
message_count: int
|
||||
flushed_at: str = field(default_factory=lambda: datetime.now(UTC).isoformat())
|
||||
tags: list[str] = field(default_factory=list)
|
||||
|
||||
@property
|
||||
def key(self) -> str:
|
||||
return f"{self.chat_id}:{self.thread_id}"
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, data: dict[str, object]) -> ThreadSummary:
|
||||
"""Deserialise from a raw dict, tolerating missing optional fields."""
|
||||
return cls(
|
||||
chat_id=int(data["chat_id"]), # type: ignore[arg-type]
|
||||
thread_id=int(data["thread_id"]), # type: ignore[arg-type]
|
||||
summary=str(data["summary"]),
|
||||
message_count=int(data["message_count"]), # type: ignore[arg-type]
|
||||
flushed_at=str(data.get("flushed_at", datetime.now(UTC).isoformat())),
|
||||
tags=list(data.get("tags", [])), # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
def format_for_telegram(self) -> str:
|
||||
"""Return a concise Telegram-formatted recall card."""
|
||||
flushed = self.flushed_at[:19].replace("T", " ")
|
||||
header = (
|
||||
f"\U0001f9e0 *Thread summary* (thread `{self.thread_id}`)\n"
|
||||
f"_Flushed: {flushed} UTC \u2014 {self.message_count} messages_"
|
||||
)
|
||||
tag_line = f"\U0001f3f7 Tags: {', '.join(self.tags)}" if self.tags else ""
|
||||
parts = [header, tag_line, self.summary] if tag_line else [header, self.summary]
|
||||
return "\n\n".join(parts)
|
||||
|
||||
|
||||
class ThreadMemoryStore:
|
||||
"""JSON-backed knowledge-base store for flushed thread summaries.
|
||||
|
||||
Summaries are indexed by keyword tags so they can be recalled contextually
|
||||
(via :meth:`search`) without being kept permanently in the LLM context.
|
||||
Each flush overwrites the previous summary for the same (chat_id, thread_id) pair.
|
||||
"""
|
||||
|
||||
def __init__(self, path: str | Path = "thread_memory.json") -> None:
|
||||
self._path = Path(path)
|
||||
self._data: dict[str, dict[str, object]] = self._load()
|
||||
|
||||
def _load(self) -> dict[str, dict[str, object]]:
|
||||
if self._path.exists():
|
||||
try:
|
||||
raw = json.loads(self._path.read_text(encoding="utf-8"))
|
||||
if isinstance(raw, dict):
|
||||
return raw # type: ignore[return-value]
|
||||
except (json.JSONDecodeError, OSError):
|
||||
logger.warning(
|
||||
"Could not read thread memory store at %s; starting fresh", self._path
|
||||
)
|
||||
return {}
|
||||
|
||||
def _save(self) -> None:
|
||||
try:
|
||||
self._path.write_text(
|
||||
json.dumps(self._data, indent=2, ensure_ascii=False), encoding="utf-8"
|
||||
)
|
||||
except OSError:
|
||||
logger.exception("Failed to write thread memory store to %s", self._path)
|
||||
|
||||
def save(self, summary: ThreadSummary) -> None:
|
||||
"""Persist a thread summary, replacing any previous entry for this thread."""
|
||||
self._data[summary.key] = asdict(summary)
|
||||
self._save()
|
||||
|
||||
def get(self, chat_id: int, thread_id: int) -> ThreadSummary | None:
|
||||
"""Return the stored summary for a thread, or None if not found."""
|
||||
raw = self._data.get(f"{chat_id}:{thread_id}")
|
||||
if raw is None:
|
||||
return None
|
||||
return ThreadSummary.from_dict(raw)
|
||||
|
||||
def all(self) -> list[ThreadSummary]:
|
||||
"""Return all stored summaries, newest first."""
|
||||
entries = [ThreadSummary.from_dict(v) for v in self._data.values()]
|
||||
return sorted(entries, key=lambda s: s.flushed_at, reverse=True)
|
||||
|
||||
def search(self, query: str) -> list[ThreadSummary]:
|
||||
"""Return summaries whose tags overlap with words in *query*, newest first.
|
||||
|
||||
The match is case-insensitive and word-based. Summaries without tags
|
||||
are not returned even if the query is broad.
|
||||
"""
|
||||
query_words = {w.lower() for w in query.split() if w}
|
||||
if not query_words:
|
||||
return []
|
||||
results = [
|
||||
ThreadSummary.from_dict(v)
|
||||
for v in self._data.values()
|
||||
if query_words & {t.lower() for t in v.get("tags", [])} # type: ignore[union-attr]
|
||||
]
|
||||
return sorted(results, key=lambda s: s.flushed_at, reverse=True)
|
||||
0
steward/proposals/__init__.py
Normal file
0
steward/proposals/__init__.py
Normal file
106
steward/proposals/generator.py
Normal file
106
steward/proposals/generator.py
Normal file
@ -0,0 +1,106 @@
|
||||
"""Proposal data model and generator.
|
||||
|
||||
This module implements a minimal "generate proposal" workflow:
|
||||
|
||||
1. Fetch data from a target API endpoint.
|
||||
2. Ask the LLM to analyse the data and produce a structured proposal.
|
||||
3. Return the proposal for delivery (e.g. via Telegram).
|
||||
|
||||
The Proposal dataclass is intentionally simple for the MVP – it captures the
|
||||
fields described in the manifesto (why, confidence, evidence, expected outcome,
|
||||
rollback strategy) without any persistence layer yet.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import UTC, datetime
|
||||
|
||||
import httpx
|
||||
|
||||
from steward.config import Settings
|
||||
from steward.llm.client import LLMClient
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_ANALYSIS_SYSTEM_PROMPT = (
|
||||
"You are Steward, an AI operations platform. "
|
||||
"You have been given raw data from an internal API. "
|
||||
"Analyse the data and produce a concise proposal in the following format:\n\n"
|
||||
"**Summary:** <one-sentence summary>\n"
|
||||
"**Why:** <reason this proposal matters>\n"
|
||||
"**Evidence:** <key data points from the API response>\n"
|
||||
"**Expected outcome:** <what will improve if adopted>\n"
|
||||
"**Rollback strategy:** <how to undo if things go wrong>\n"
|
||||
"**Confidence:** <Low | Medium | High> - <brief justification>\n\n"
|
||||
"Be conservative. If the data is healthy and no action is needed, say so explicitly."
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class Proposal:
|
||||
"""A structured action proposal generated by Steward."""
|
||||
|
||||
title: str
|
||||
body: str
|
||||
source_url: str
|
||||
generated_at: datetime = field(default_factory=lambda: datetime.now(UTC))
|
||||
raw_data: str = ""
|
||||
|
||||
def format_for_telegram(self) -> str:
|
||||
"""Return a Markdown-formatted string suitable for a Telegram message."""
|
||||
ts = self.generated_at.strftime("%Y-%m-%d %H:%M UTC")
|
||||
return (
|
||||
f"\U0001f50d *Steward Proposal*\n"
|
||||
f"_{ts}_\n\n"
|
||||
f"*Source:* `{self.source_url}`\n\n"
|
||||
f"{self.body}"
|
||||
)
|
||||
|
||||
|
||||
class ProposalGenerator:
|
||||
"""Fetches data from a target API and generates a proposal via the LLM."""
|
||||
|
||||
def __init__(self, settings: Settings, llm: LLMClient) -> None:
|
||||
self._settings = settings
|
||||
self._llm = llm
|
||||
|
||||
async def run(self) -> Proposal | None:
|
||||
"""Fetch the target API and return a Proposal, or None on error."""
|
||||
url = self._settings.analysis_target_url
|
||||
if not url:
|
||||
logger.warning("analysis_target_url is not configured - skipping proposal generation")
|
||||
return None
|
||||
|
||||
raw = await self._fetch(url)
|
||||
if raw is None:
|
||||
return None
|
||||
|
||||
body = await self._llm.chat(
|
||||
f"Here is the API response from {url}:\n\n{raw}",
|
||||
system_prompt=_ANALYSIS_SYSTEM_PROMPT,
|
||||
)
|
||||
|
||||
return Proposal(
|
||||
title="Daily Analysis",
|
||||
body=body,
|
||||
source_url=url,
|
||||
raw_data=raw,
|
||||
)
|
||||
|
||||
async def _fetch(self, url: str) -> str | None:
|
||||
"""Fetch the URL and return the response body as text."""
|
||||
headers: dict[str, str] = {}
|
||||
api_key = self._settings.analysis_target_api_key
|
||||
if api_key:
|
||||
headers["Authorization"] = "Bearer " + api_key
|
||||
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=30) as client:
|
||||
response = await client.get(url, headers=headers)
|
||||
response.raise_for_status()
|
||||
return response.text
|
||||
except httpx.HTTPError as exc:
|
||||
logger.error("Failed to fetch %s: %s", url, exc)
|
||||
return None
|
||||
1
steward/tools/__init__.py
Normal file
1
steward/tools/__init__.py
Normal file
@ -0,0 +1 @@
|
||||
"""MCP-compatible OpenAPI tool calling support."""
|
||||
259
steward/tools/client.py
Normal file
259
steward/tools/client.py
Normal file
@ -0,0 +1,259 @@
|
||||
"""OpenAPI-based tool client for MCP-compatible tool servers.
|
||||
|
||||
Compatible with open-webui/openapi-servers (https://github.com/open-webui/openapi-servers).
|
||||
Fetches the server's OpenAPI spec, converts operations to OpenAI function-calling
|
||||
definitions, and executes tool calls by making HTTP requests back to the server.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class ToolClient:
|
||||
"""Client for OpenAPI-compatible tool servers.
|
||||
|
||||
Usage::
|
||||
|
||||
client = ToolClient(base_url="http://localhost:8000", api_key="…")
|
||||
tools = await client.get_tools() # OpenAI tool definitions
|
||||
result = await client.call("my_op", {}) # execute a tool call
|
||||
"""
|
||||
|
||||
def __init__(self, base_url: str, api_key: str = "", timeout: float = 30.0) -> None:
|
||||
self._base_url = base_url.rstrip("/")
|
||||
self._api_key = api_key
|
||||
self._timeout = timeout
|
||||
self._spec: dict[str, Any] | None = None
|
||||
self._tools: list[dict[str, Any]] | None = None
|
||||
|
||||
@property
|
||||
def base_url(self) -> str:
|
||||
return self._base_url
|
||||
|
||||
def _headers(self) -> dict[str, str]:
|
||||
headers: dict[str, str] = {"Accept": "application/json"}
|
||||
if self._api_key:
|
||||
headers["Authorization"] = "Bearer " + self._api_key.strip()
|
||||
return headers
|
||||
|
||||
async def get_spec(self) -> dict[str, Any]:
|
||||
"""Fetch (and cache) the OpenAPI spec from ``/openapi.json``."""
|
||||
if self._spec is not None:
|
||||
return self._spec
|
||||
async with httpx.AsyncClient(timeout=self._timeout) as http:
|
||||
resp = await http.get(
|
||||
f"{self._base_url}/openapi.json",
|
||||
headers=self._headers(),
|
||||
)
|
||||
resp.raise_for_status()
|
||||
self._spec = resp.json()
|
||||
path_count = len(self._spec.get("paths", {}))
|
||||
logger.info("Loaded OpenAPI spec from %s (%d paths)", self._base_url, path_count)
|
||||
return self._spec
|
||||
|
||||
async def get_tools(self) -> list[dict[str, Any]]:
|
||||
"""Return OpenAI function-calling tool definitions from the OpenAPI spec."""
|
||||
if self._tools is not None:
|
||||
return self._tools
|
||||
spec = await self.get_spec()
|
||||
self._tools = _spec_to_openai_tools(spec)
|
||||
logger.info("Registered %d tools from %s", len(self._tools), self._base_url)
|
||||
return self._tools
|
||||
|
||||
async def call(self, tool_name: str, arguments: dict[str, Any]) -> str:
|
||||
"""Execute a named tool call and return the result as a string.
|
||||
|
||||
Resolves *tool_name* to a path+method in the cached OpenAPI spec,
|
||||
partitions *arguments* into path params, query params, and request
|
||||
body, then makes the HTTP request.
|
||||
"""
|
||||
spec = await self.get_spec()
|
||||
path, method, path_params, query_params, body = _resolve_operation(
|
||||
spec, tool_name, arguments
|
||||
)
|
||||
|
||||
url = self._base_url + path
|
||||
for key, value in path_params.items():
|
||||
url = url.replace(f"{{{key}}}", str(value))
|
||||
|
||||
async with httpx.AsyncClient(timeout=self._timeout) as http:
|
||||
resp = await http.request(
|
||||
method=method.upper(),
|
||||
url=url,
|
||||
headers=self._headers(),
|
||||
params=query_params if query_params else None,
|
||||
json=body if body else None,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
content_type = resp.headers.get("content-type", "")
|
||||
if "application/json" in content_type:
|
||||
try:
|
||||
return json.dumps(resp.json(), indent=2)
|
||||
except ValueError:
|
||||
pass
|
||||
return resp.text
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# OpenAPI → OpenAI tool-definition helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _resolve_ref(
|
||||
schema: dict[str, Any], components: dict[str, Any]
|
||||
) -> dict[str, Any]:
|
||||
"""Recursively resolve a ``$ref`` inside an OpenAPI schema."""
|
||||
if "$ref" not in schema:
|
||||
return schema
|
||||
ref: str = schema["$ref"] # e.g. "#/components/schemas/Foo"
|
||||
parts = ref.lstrip("#/").split("/")
|
||||
# parts = ["components", "schemas", "Foo"]
|
||||
obj: Any = {"components": components}
|
||||
for part in parts:
|
||||
if not isinstance(obj, dict):
|
||||
return {}
|
||||
obj = obj.get(part, {})
|
||||
return _schema_to_json_schema(obj, components)
|
||||
|
||||
|
||||
def _schema_to_json_schema(
|
||||
schema: dict[str, Any], components: dict[str, Any] | None = None
|
||||
) -> dict[str, Any]:
|
||||
"""Convert an OpenAPI schema object to a JSON Schema-compatible dict."""
|
||||
if not schema:
|
||||
return {}
|
||||
resolved = _resolve_ref(schema, components or {}) if "$ref" in schema else schema
|
||||
result: dict[str, Any] = {}
|
||||
for key in ("type", "description", "enum", "format", "default", "minimum", "maximum"):
|
||||
if key in resolved:
|
||||
result[key] = resolved[key]
|
||||
if "items" in resolved:
|
||||
result["items"] = _schema_to_json_schema(resolved["items"], components)
|
||||
if "properties" in resolved:
|
||||
result["properties"] = {
|
||||
k: _schema_to_json_schema(v, components)
|
||||
for k, v in resolved["properties"].items()
|
||||
}
|
||||
if "required" in resolved:
|
||||
result["required"] = resolved["required"]
|
||||
return result
|
||||
|
||||
|
||||
def _operation_id(method: str, path: str, op: dict[str, Any]) -> str:
|
||||
"""Return an operation ID, generating one from method+path if absent."""
|
||||
if op.get("operationId"):
|
||||
return str(op["operationId"])
|
||||
slug = path.strip("/").replace("/", "_").replace("{", "").replace("}", "")
|
||||
return f"{method}_{slug}" if slug else method
|
||||
|
||||
|
||||
def _spec_to_openai_tools(spec: dict[str, Any]) -> list[dict[str, Any]]:
|
||||
"""Convert an OpenAPI 3.x spec to a list of OpenAI function-calling tool defs."""
|
||||
tools: list[dict[str, Any]] = []
|
||||
paths: dict[str, Any] = spec.get("paths", {})
|
||||
components: dict[str, Any] = spec.get("components", {})
|
||||
|
||||
for path, path_item in paths.items():
|
||||
for http_method in ("get", "post", "put", "patch", "delete"):
|
||||
op: dict[str, Any] | None = path_item.get(http_method)
|
||||
if op is None:
|
||||
continue
|
||||
|
||||
op_id = _operation_id(http_method, path, op)
|
||||
description: str = op.get("summary") or op.get("description") or ""
|
||||
|
||||
properties: dict[str, Any] = {}
|
||||
required: list[str] = []
|
||||
|
||||
# URL / query parameters
|
||||
for param in op.get("parameters", []):
|
||||
name: str = param["name"]
|
||||
schema = _schema_to_json_schema(
|
||||
param.get("schema", {"type": "string"}), components
|
||||
)
|
||||
if param.get("description"):
|
||||
schema["description"] = param["description"]
|
||||
properties[name] = schema
|
||||
if param.get("required"):
|
||||
required.append(name)
|
||||
|
||||
# Request body – flatten top-level object properties
|
||||
rb: dict[str, Any] = op.get("requestBody", {})
|
||||
if rb:
|
||||
json_content = rb.get("content", {}).get("application/json", {})
|
||||
body_schema = _schema_to_json_schema(
|
||||
json_content.get("schema", {}), components
|
||||
)
|
||||
if body_schema.get("type") == "object":
|
||||
for prop_name, prop_schema in body_schema.get("properties", {}).items():
|
||||
properties[prop_name] = prop_schema
|
||||
if rb.get("required", False):
|
||||
required.extend(body_schema.get("required", []))
|
||||
|
||||
tool: dict[str, Any] = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": op_id,
|
||||
"description": description,
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": properties,
|
||||
},
|
||||
},
|
||||
}
|
||||
if required:
|
||||
tool["function"]["parameters"]["required"] = required
|
||||
|
||||
tools.append(tool)
|
||||
|
||||
return tools
|
||||
|
||||
|
||||
def _resolve_operation(
|
||||
spec: dict[str, Any],
|
||||
operation_id: str,
|
||||
arguments: dict[str, Any],
|
||||
) -> tuple[str, str, dict[str, Any], dict[str, Any], dict[str, Any]]:
|
||||
"""Find path+method for *operation_id* and split *arguments* into parts.
|
||||
|
||||
Returns ``(path, method, path_params, query_params, body)``.
|
||||
|
||||
* ``path_params`` – values that replace ``{placeholders}`` in the URL path.
|
||||
* ``query_params`` – values appended as query string.
|
||||
* ``body`` – any remaining arguments sent as the JSON request body.
|
||||
"""
|
||||
paths: dict[str, Any] = spec.get("paths", {})
|
||||
|
||||
for path, path_item in paths.items():
|
||||
for http_method in ("get", "post", "put", "patch", "delete"):
|
||||
op: dict[str, Any] | None = path_item.get(http_method)
|
||||
if op is None:
|
||||
continue
|
||||
if _operation_id(http_method, path, op) != operation_id:
|
||||
continue
|
||||
|
||||
path_params: dict[str, Any] = {}
|
||||
query_params: dict[str, Any] = {}
|
||||
declared_params: set[str] = set()
|
||||
|
||||
for param in op.get("parameters", []):
|
||||
name = param["name"]
|
||||
declared_params.add(name)
|
||||
if name not in arguments:
|
||||
continue
|
||||
if param.get("in") == "path":
|
||||
path_params[name] = arguments[name]
|
||||
else:
|
||||
query_params[name] = arguments[name]
|
||||
|
||||
body = {k: v for k, v in arguments.items() if k not in declared_params}
|
||||
return path, http_method, path_params, query_params, body
|
||||
|
||||
raise ValueError(f"Operation '{operation_id}' not found in spec")
|
||||
0
tests/__init__.py
Normal file
0
tests/__init__.py
Normal file
184
tests/test_bot.py
Normal file
184
tests/test_bot.py
Normal file
@ -0,0 +1,184 @@
|
||||
"""Tests for steward.bot.telegram (handler logic)."""
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from telegram import Chat, Message, Update, User
|
||||
from telegram.ext import CallbackContext
|
||||
|
||||
from steward.bot.telegram import (
|
||||
_history,
|
||||
_is_allowed,
|
||||
_thread_history,
|
||||
clear_handler,
|
||||
message_handler,
|
||||
start_handler,
|
||||
)
|
||||
from steward.config import Settings
|
||||
from steward.llm.client import LLMClient
|
||||
from steward.memory.thread_store import ThreadMemoryStore
|
||||
|
||||
|
||||
def _make_settings(**kwargs) -> Settings:
|
||||
defaults = dict(
|
||||
telegram_bot_token="test-token",
|
||||
openai_api_key="test-key",
|
||||
)
|
||||
defaults.update(kwargs)
|
||||
return Settings(**defaults)
|
||||
|
||||
|
||||
def _make_update(
|
||||
user_id: int = 12345,
|
||||
text: str = "hello",
|
||||
chat_id: int | None = None,
|
||||
thread_id: int | None = None,
|
||||
) -> Update:
|
||||
user = MagicMock(spec=User)
|
||||
user.id = user_id
|
||||
|
||||
message = MagicMock(spec=Message)
|
||||
message.text = text
|
||||
message.reply_text = AsyncMock()
|
||||
message.message_thread_id = thread_id
|
||||
|
||||
chat = MagicMock(spec=Chat)
|
||||
chat.id = chat_id if chat_id is not None else user_id
|
||||
|
||||
update = MagicMock(spec=Update)
|
||||
update.effective_user = user
|
||||
update.effective_chat = chat
|
||||
update.message = message
|
||||
return update
|
||||
|
||||
|
||||
def _make_context(
|
||||
settings: Settings,
|
||||
llm: LLMClient | None = None,
|
||||
store: ThreadMemoryStore | None = None,
|
||||
) -> CallbackContext: # type: ignore[type-arg]
|
||||
ctx = MagicMock(spec=CallbackContext)
|
||||
if store is None:
|
||||
mock_store = MagicMock(spec=ThreadMemoryStore)
|
||||
mock_store.search.return_value = []
|
||||
store = mock_store
|
||||
ctx.bot_data = {
|
||||
"settings": settings,
|
||||
"llm": llm,
|
||||
"thread_store": store,
|
||||
}
|
||||
ctx.args = []
|
||||
return ctx
|
||||
|
||||
|
||||
class TestIsAllowed:
|
||||
def test_no_allowlist_allows_everyone(self):
|
||||
settings = _make_settings(telegram_allowed_user_ids=[])
|
||||
assert _is_allowed(999, settings) is True
|
||||
|
||||
def test_allowlist_accepts_known_user(self):
|
||||
settings = _make_settings(telegram_allowed_user_ids=[1, 2, 3])
|
||||
assert _is_allowed(2, settings) is True
|
||||
|
||||
def test_allowlist_rejects_unknown_user(self):
|
||||
settings = _make_settings(telegram_allowed_user_ids=[1, 2, 3])
|
||||
assert _is_allowed(999, settings) is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_start_handler_replies(monkeypatch):
|
||||
settings = _make_settings()
|
||||
update = _make_update()
|
||||
ctx = _make_context(settings)
|
||||
|
||||
await start_handler(update, ctx)
|
||||
|
||||
update.message.reply_text.assert_awaited_once()
|
||||
text = update.message.reply_text.call_args.args[0]
|
||||
assert "Steward" in text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_start_handler_ignores_disallowed_user():
|
||||
settings = _make_settings(telegram_allowed_user_ids=[9999])
|
||||
update = _make_update(user_id=1111)
|
||||
ctx = _make_context(settings)
|
||||
|
||||
await start_handler(update, ctx)
|
||||
|
||||
update.message.reply_text.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_clear_handler_clears_user_history():
|
||||
"""clear_handler without a thread clears per-user history."""
|
||||
settings = _make_settings()
|
||||
user_id = 42
|
||||
_history[user_id] = [{"role": "user", "content": "old msg"}]
|
||||
|
||||
update = _make_update(user_id=user_id) # no thread_id
|
||||
ctx = _make_context(settings)
|
||||
|
||||
await clear_handler(update, ctx)
|
||||
|
||||
assert _history[user_id] == []
|
||||
update.message.reply_text.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_clear_handler_clears_thread_history():
|
||||
"""clear_handler inside a thread clears that thread's history."""
|
||||
settings = _make_settings()
|
||||
key = (100, 7)
|
||||
_thread_history[key] = [{"role": "user", "content": "thread msg"}]
|
||||
|
||||
update = _make_update(user_id=1, chat_id=100, thread_id=7)
|
||||
ctx = _make_context(settings)
|
||||
|
||||
await clear_handler(update, ctx)
|
||||
|
||||
assert _thread_history[key] == []
|
||||
update.message.reply_text.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_message_handler_calls_llm_and_replies():
|
||||
settings = _make_settings()
|
||||
mock_llm = MagicMock(spec=LLMClient)
|
||||
mock_llm.chat = AsyncMock(return_value="LLM response")
|
||||
|
||||
update = _make_update(user_id=77, text="What is the weather?")
|
||||
ctx = _make_context(settings, llm=mock_llm)
|
||||
|
||||
_history[77].clear()
|
||||
await message_handler(update, ctx)
|
||||
|
||||
mock_llm.chat.assert_awaited_once()
|
||||
update.message.reply_text.assert_awaited_once()
|
||||
assert len(_history[77]) == 2
|
||||
assert _history[77][0]["role"] == "user"
|
||||
assert _history[77][1]["role"] == "assistant"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_message_handler_non_thread_trims_history():
|
||||
"""Non-thread history is trimmed to _MAX_HISTORY turns."""
|
||||
from steward.bot.telegram import _MAX_HISTORY
|
||||
|
||||
settings = _make_settings()
|
||||
mock_llm = MagicMock(spec=LLMClient)
|
||||
mock_llm.chat = AsyncMock(return_value="reply")
|
||||
|
||||
user_id = 200
|
||||
# Pre-fill exactly at the limit
|
||||
_history[user_id] = [
|
||||
{"role": "user" if i % 2 == 0 else "assistant", "content": f"msg{i}"}
|
||||
for i in range(_MAX_HISTORY * 2)
|
||||
]
|
||||
|
||||
update = _make_update(user_id=user_id, text="new question")
|
||||
ctx = _make_context(settings, llm=mock_llm)
|
||||
await message_handler(update, ctx)
|
||||
|
||||
# After appending 2 new messages, should trim back to _MAX_HISTORY * 2
|
||||
assert len(_history[user_id]) == _MAX_HISTORY * 2
|
||||
92
tests/test_llm_client.py
Normal file
92
tests/test_llm_client.py
Normal file
@ -0,0 +1,92 @@
|
||||
"""Tests for steward.llm.client."""
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from steward.config import Settings
|
||||
from steward.llm.client import LLMClient
|
||||
|
||||
|
||||
def _make_settings(**kwargs) -> Settings:
|
||||
defaults = dict(
|
||||
telegram_bot_token="test-token",
|
||||
openai_api_key="test-key",
|
||||
openai_model="gpt-4o",
|
||||
openai_system_prompt="You are Steward.",
|
||||
)
|
||||
defaults.update(kwargs)
|
||||
return Settings(**defaults)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def settings():
|
||||
return _make_settings()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def llm_client(settings):
|
||||
return LLMClient(settings)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_sends_correct_messages(llm_client):
|
||||
"""chat() should prepend the system prompt and append the user message."""
|
||||
mock_response = MagicMock()
|
||||
mock_response.choices[0].message.content = "Hello from the LLM!"
|
||||
|
||||
with patch.object(
|
||||
llm_client._client.chat.completions,
|
||||
"create",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_response,
|
||||
) as mock_create:
|
||||
result = await llm_client.chat("Hi there")
|
||||
|
||||
assert result == "Hello from the LLM!"
|
||||
call_kwargs = mock_create.call_args.kwargs
|
||||
messages = call_kwargs["messages"]
|
||||
assert messages[0]["role"] == "system"
|
||||
assert messages[-1]["role"] == "user"
|
||||
assert messages[-1]["content"] == "Hi there"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_with_history(llm_client):
|
||||
"""chat() should include history between system and user messages."""
|
||||
history = [
|
||||
{"role": "user", "content": "previous question"},
|
||||
{"role": "assistant", "content": "previous answer"},
|
||||
]
|
||||
mock_response = MagicMock()
|
||||
mock_response.choices[0].message.content = "reply"
|
||||
|
||||
with patch.object(
|
||||
llm_client._client.chat.completions,
|
||||
"create",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_response,
|
||||
) as mock_create:
|
||||
await llm_client.chat("new question", history=history)
|
||||
|
||||
messages = mock_create.call_args.kwargs["messages"]
|
||||
roles = [m["role"] for m in messages]
|
||||
assert roles == ["system", "user", "assistant", "user"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_custom_system_prompt(llm_client):
|
||||
"""chat() should use a custom system prompt when provided."""
|
||||
mock_response = MagicMock()
|
||||
mock_response.choices[0].message.content = "ok"
|
||||
|
||||
with patch.object(
|
||||
llm_client._client.chat.completions,
|
||||
"create",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_response,
|
||||
) as mock_create:
|
||||
await llm_client.chat("msg", system_prompt="Custom prompt")
|
||||
|
||||
messages = mock_create.call_args.kwargs["messages"]
|
||||
assert messages[0]["content"] == "Custom prompt"
|
||||
106
tests/test_proposals.py
Normal file
106
tests/test_proposals.py
Normal file
@ -0,0 +1,106 @@
|
||||
"""Tests for steward.proposals.generator."""
|
||||
|
||||
from datetime import UTC
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import respx
|
||||
|
||||
from steward.config import Settings
|
||||
from steward.llm.client import LLMClient
|
||||
from steward.proposals.generator import Proposal, ProposalGenerator
|
||||
|
||||
|
||||
def _make_settings(**kwargs) -> Settings:
|
||||
defaults = dict(
|
||||
openai_api_key="test-key",
|
||||
analysis_target_url="https://example.com/api/status",
|
||||
analysis_target_api_key="secret",
|
||||
)
|
||||
defaults.update(kwargs)
|
||||
return Settings(**defaults)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def settings():
|
||||
return _make_settings()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_llm():
|
||||
llm = MagicMock(spec=LLMClient)
|
||||
llm.chat = AsyncMock(
|
||||
return_value="**Summary:** Everything is healthy.\n**Why:** No issues detected."
|
||||
)
|
||||
return llm
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_generates_proposal(settings, mock_llm):
|
||||
"""run() should fetch the URL and return a Proposal."""
|
||||
with respx.mock:
|
||||
respx.get("https://example.com/api/status").mock(
|
||||
return_value=httpx.Response(200, text='{"status": "ok"}')
|
||||
)
|
||||
generator = ProposalGenerator(settings, mock_llm)
|
||||
proposal = await generator.run()
|
||||
|
||||
assert proposal is not None
|
||||
assert isinstance(proposal, Proposal)
|
||||
assert proposal.source_url == "https://example.com/api/status"
|
||||
assert proposal.raw_data == '{"status": "ok"}'
|
||||
mock_llm.chat.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_no_url_returns_none(mock_llm):
|
||||
"""run() should return None when analysis_target_url is not configured."""
|
||||
settings = _make_settings(analysis_target_url="")
|
||||
generator = ProposalGenerator(settings, mock_llm)
|
||||
result = await generator.run()
|
||||
assert result is None
|
||||
mock_llm.chat.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_http_error_returns_none(settings, mock_llm):
|
||||
"""run() should return None when the HTTP request fails."""
|
||||
with respx.mock:
|
||||
respx.get("https://example.com/api/status").mock(
|
||||
return_value=httpx.Response(500, text="Internal Server Error")
|
||||
)
|
||||
generator = ProposalGenerator(settings, mock_llm)
|
||||
proposal = await generator.run()
|
||||
|
||||
assert proposal is None
|
||||
mock_llm.chat.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_sends_bearer_token(settings, mock_llm):
|
||||
"""_fetch() should include an Authorization header when api_key is configured."""
|
||||
with respx.mock:
|
||||
route = respx.get("https://example.com/api/status").mock(
|
||||
return_value=httpx.Response(200, text="data")
|
||||
)
|
||||
generator = ProposalGenerator(settings, mock_llm)
|
||||
await generator._fetch("https://example.com/api/status")
|
||||
|
||||
request = route.calls.last.request
|
||||
assert request.headers["Authorization"].startswith("Bearer ")
|
||||
|
||||
|
||||
def test_proposal_format_for_telegram():
|
||||
"""format_for_telegram() should contain the source URL and body."""
|
||||
from datetime import datetime
|
||||
p = Proposal(
|
||||
title="Test",
|
||||
body="**Summary:** All good.",
|
||||
source_url="https://example.com",
|
||||
generated_at=datetime(2026, 1, 1, 8, 0, 0, tzinfo=UTC),
|
||||
)
|
||||
text = p.format_for_telegram()
|
||||
assert p.source_url in text
|
||||
assert "All good." in text
|
||||
assert "2026-01-01" in text
|
||||
500
tests/test_thread_memory.py
Normal file
500
tests/test_thread_memory.py
Normal file
@ -0,0 +1,500 @@
|
||||
"""Tests for thread-aware memory: ThreadMemoryStore and /flush, /recall handlers."""
|
||||
|
||||
from pathlib import Path
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from telegram import Chat, Message, Update, User
|
||||
from telegram.ext import CallbackContext
|
||||
|
||||
from steward.bot.telegram import (
|
||||
_thread_history,
|
||||
flush_handler,
|
||||
message_handler,
|
||||
recall_handler,
|
||||
)
|
||||
from steward.config import Settings
|
||||
from steward.llm.client import LLMClient
|
||||
from steward.memory.thread_store import ThreadMemoryStore, ThreadSummary
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _make_settings(**kwargs) -> Settings:
|
||||
return Settings(telegram_bot_token="tok", openai_api_key="key", **kwargs)
|
||||
|
||||
|
||||
def _make_update(
|
||||
user_id: int = 1,
|
||||
text: str = "hi",
|
||||
chat_id: int = 100,
|
||||
thread_id: int | None = None,
|
||||
) -> Update:
|
||||
user = MagicMock(spec=User)
|
||||
user.id = user_id
|
||||
|
||||
message = MagicMock(spec=Message)
|
||||
message.text = text
|
||||
message.reply_text = AsyncMock()
|
||||
message.message_thread_id = thread_id
|
||||
|
||||
chat = MagicMock(spec=Chat)
|
||||
chat.id = chat_id
|
||||
|
||||
update = MagicMock(spec=Update)
|
||||
update.effective_user = user
|
||||
update.effective_chat = chat
|
||||
update.message = message
|
||||
return update
|
||||
|
||||
|
||||
def _make_context(
|
||||
settings: Settings,
|
||||
llm: LLMClient | None = None,
|
||||
store: ThreadMemoryStore | None = None,
|
||||
) -> CallbackContext: # type: ignore[type-arg]
|
||||
ctx = MagicMock(spec=CallbackContext)
|
||||
if store is None:
|
||||
mock_store = MagicMock(spec=ThreadMemoryStore)
|
||||
mock_store.search.return_value = []
|
||||
store = mock_store
|
||||
ctx.bot_data = {
|
||||
"settings": settings,
|
||||
"llm": llm or MagicMock(spec=LLMClient),
|
||||
"thread_store": store,
|
||||
}
|
||||
ctx.args = []
|
||||
return ctx
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ThreadMemoryStore unit tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestThreadMemoryStore:
|
||||
def _store(self, tmp_path: Path) -> ThreadMemoryStore:
|
||||
return ThreadMemoryStore(tmp_path / "mem.json")
|
||||
|
||||
def test_get_missing_returns_none(self, tmp_path: Path):
|
||||
store = self._store(tmp_path)
|
||||
assert store.get(1, 2) is None
|
||||
|
||||
def test_save_and_get_roundtrip(self, tmp_path: Path):
|
||||
store = self._store(tmp_path)
|
||||
summary = ThreadSummary(chat_id=1, thread_id=42, summary="A recap.", message_count=5)
|
||||
store.save(summary)
|
||||
|
||||
retrieved = store.get(1, 42)
|
||||
assert retrieved is not None
|
||||
assert retrieved.summary == "A recap."
|
||||
assert retrieved.message_count == 5
|
||||
|
||||
def test_save_overwrites_previous_entry(self, tmp_path: Path):
|
||||
store = self._store(tmp_path)
|
||||
store.save(ThreadSummary(chat_id=1, thread_id=1, summary="old", message_count=2))
|
||||
store.save(ThreadSummary(chat_id=1, thread_id=1, summary="new", message_count=4))
|
||||
assert store.get(1, 1).summary == "new" # type: ignore[union-attr]
|
||||
|
||||
def test_persisted_to_disk(self, tmp_path: Path):
|
||||
path = tmp_path / "mem.json"
|
||||
store = ThreadMemoryStore(path)
|
||||
store.save(ThreadSummary(chat_id=5, thread_id=9, summary="saved", message_count=1))
|
||||
|
||||
# Load a fresh store from the same file
|
||||
store2 = ThreadMemoryStore(path)
|
||||
assert store2.get(5, 9) is not None
|
||||
|
||||
def test_all_returns_newest_first(self, tmp_path: Path):
|
||||
store = self._store(tmp_path)
|
||||
store.save(ThreadSummary(chat_id=1, thread_id=1, summary="first", message_count=1,
|
||||
flushed_at="2026-01-01T00:00:00+00:00"))
|
||||
store.save(ThreadSummary(chat_id=1, thread_id=2, summary="second", message_count=1,
|
||||
flushed_at="2026-06-01T00:00:00+00:00"))
|
||||
results = store.all()
|
||||
assert results[0].summary == "second"
|
||||
assert results[1].summary == "first"
|
||||
|
||||
def test_corrupt_file_starts_fresh(self, tmp_path: Path):
|
||||
path = tmp_path / "mem.json"
|
||||
path.write_text("not-json", encoding="utf-8")
|
||||
store = ThreadMemoryStore(path)
|
||||
assert store.all() == []
|
||||
|
||||
def test_format_for_telegram_contains_thread_id(self, tmp_path: Path):
|
||||
s = ThreadSummary(chat_id=1, thread_id=77, summary="recap", message_count=3)
|
||||
text = s.format_for_telegram()
|
||||
assert "77" in text
|
||||
assert "recap" in text
|
||||
|
||||
def test_format_for_telegram_shows_tags(self, tmp_path: Path):
|
||||
s = ThreadSummary(
|
||||
chat_id=1, thread_id=77, summary="recap", message_count=3,
|
||||
tags=["api", "auth"],
|
||||
)
|
||||
text = s.format_for_telegram()
|
||||
assert "api" in text
|
||||
assert "auth" in text
|
||||
|
||||
def test_from_dict_tolerates_missing_tags(self, tmp_path: Path):
|
||||
"""from_dict must handle legacy entries that predate the tags field."""
|
||||
raw = {"chat_id": 1, "thread_id": 2, "summary": "old", "message_count": 3,
|
||||
"flushed_at": "2026-01-01T00:00:00+00:00"}
|
||||
s = ThreadSummary.from_dict(raw)
|
||||
assert s.tags == []
|
||||
|
||||
def test_search_returns_matching_summaries(self, tmp_path: Path):
|
||||
store = self._store(tmp_path)
|
||||
store.save(ThreadSummary(chat_id=1, thread_id=1, summary="API work", message_count=2,
|
||||
tags=["api", "design", "auth"]))
|
||||
store.save(ThreadSummary(chat_id=1, thread_id=2, summary="Database work", message_count=2,
|
||||
tags=["database", "schema"]))
|
||||
results = store.search("api")
|
||||
assert len(results) == 1
|
||||
assert results[0].thread_id == 1
|
||||
|
||||
def test_search_no_match_returns_empty(self, tmp_path: Path):
|
||||
store = self._store(tmp_path)
|
||||
store.save(ThreadSummary(chat_id=1, thread_id=1, summary="recap", message_count=1,
|
||||
tags=["database"]))
|
||||
assert store.search("deployment") == []
|
||||
|
||||
def test_search_empty_query_returns_empty(self, tmp_path: Path):
|
||||
store = self._store(tmp_path)
|
||||
store.save(ThreadSummary(chat_id=1, thread_id=1, summary="recap", message_count=1,
|
||||
tags=["database"]))
|
||||
assert store.search("") == []
|
||||
|
||||
def test_search_case_insensitive(self, tmp_path: Path):
|
||||
store = self._store(tmp_path)
|
||||
store.save(ThreadSummary(chat_id=1, thread_id=1, summary="recap", message_count=1,
|
||||
tags=["API"]))
|
||||
assert len(store.search("api")) == 1
|
||||
|
||||
def test_search_skips_untagged_summaries(self, tmp_path: Path):
|
||||
store = self._store(tmp_path)
|
||||
store.save(ThreadSummary(chat_id=1, thread_id=1, summary="recap", message_count=1,
|
||||
tags=[]))
|
||||
assert store.search("api") == []
|
||||
|
||||
def test_legacy_store_roundtrip(self, tmp_path: Path):
|
||||
"""Summaries without tags survive a save/load cycle via from_dict."""
|
||||
path = tmp_path / "mem.json"
|
||||
import json as _json
|
||||
path.write_text(
|
||||
_json.dumps({
|
||||
"1:1": {"chat_id": 1, "thread_id": 1, "summary": "old",
|
||||
"message_count": 1, "flushed_at": "2026-01-01T00:00:00+00:00"}
|
||||
}),
|
||||
encoding="utf-8",
|
||||
)
|
||||
store = ThreadMemoryStore(path)
|
||||
s = store.get(1, 1)
|
||||
assert s is not None
|
||||
assert s.tags == []
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Thread-aware message_handler tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_thread_message_stored_in_thread_history():
|
||||
"""Messages in a thread go to _thread_history, not _history."""
|
||||
from steward.bot.telegram import _history
|
||||
|
||||
settings = _make_settings()
|
||||
mock_llm = MagicMock(spec=LLMClient)
|
||||
mock_llm.chat = AsyncMock(return_value="thread reply")
|
||||
|
||||
key = (100, 55)
|
||||
_thread_history[key].clear()
|
||||
|
||||
update = _make_update(user_id=1, chat_id=100, thread_id=55, text="thread message")
|
||||
ctx = _make_context(settings, llm=mock_llm)
|
||||
await message_handler(update, ctx)
|
||||
|
||||
assert len(_thread_history[key]) == 2
|
||||
assert _thread_history[key][0]["role"] == "user"
|
||||
# Regular user history untouched
|
||||
assert len(_history[1]) == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_thread_history_is_unbounded():
|
||||
"""Thread history never gets trimmed regardless of how many turns there are."""
|
||||
from steward.bot.telegram import _MAX_HISTORY
|
||||
|
||||
settings = _make_settings()
|
||||
mock_llm = MagicMock(spec=LLMClient)
|
||||
mock_llm.chat = AsyncMock(return_value="reply")
|
||||
|
||||
key = (200, 66)
|
||||
# Pre-fill well beyond the cap used for non-thread history
|
||||
_thread_history[key] = [
|
||||
{"role": "user" if i % 2 == 0 else "assistant", "content": f"m{i}"}
|
||||
for i in range(_MAX_HISTORY * 4) # 4× the normal cap
|
||||
]
|
||||
prior_len = len(_thread_history[key])
|
||||
|
||||
update = _make_update(user_id=2, chat_id=200, thread_id=66, text="more")
|
||||
ctx = _make_context(settings, llm=mock_llm)
|
||||
await message_handler(update, ctx)
|
||||
|
||||
# Should have grown by exactly 2 (user + assistant), never trimmed
|
||||
assert len(_thread_history[key]) == prior_len + 2
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# /flush handler tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_flush_outside_thread_warns():
|
||||
"""flush_handler outside a thread should warn the user."""
|
||||
settings = _make_settings()
|
||||
store = MagicMock(spec=ThreadMemoryStore)
|
||||
mock_llm = MagicMock(spec=LLMClient)
|
||||
mock_llm.chat = AsyncMock(return_value="summary")
|
||||
|
||||
update = _make_update(thread_id=None) # no thread
|
||||
ctx = _make_context(settings, llm=mock_llm, store=store)
|
||||
await flush_handler(update, ctx)
|
||||
|
||||
store.save.assert_not_called()
|
||||
update.message.reply_text.assert_awaited_once()
|
||||
text = update.message.reply_text.call_args.args[0]
|
||||
assert "flush" in text.lower() or "thread" in text.lower()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_flush_empty_thread_warns():
|
||||
"""flush_handler with no history should tell the user there is nothing to flush."""
|
||||
settings = _make_settings()
|
||||
store = MagicMock(spec=ThreadMemoryStore)
|
||||
mock_llm = MagicMock(spec=LLMClient)
|
||||
mock_llm.chat = AsyncMock(return_value="summary")
|
||||
|
||||
key = (300, 88)
|
||||
_thread_history[key].clear()
|
||||
|
||||
update = _make_update(chat_id=300, thread_id=88)
|
||||
ctx = _make_context(settings, llm=mock_llm, store=store)
|
||||
await flush_handler(update, ctx)
|
||||
|
||||
store.save.assert_not_called()
|
||||
mock_llm.chat.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_flush_summarises_stores_and_compresses():
|
||||
"""flush_handler should summarise, persist, and compress in-memory history."""
|
||||
settings = _make_settings()
|
||||
store = MagicMock(spec=ThreadMemoryStore)
|
||||
mock_llm = MagicMock(spec=LLMClient)
|
||||
# First call → summary, second call → tags
|
||||
mock_llm.chat = AsyncMock(side_effect=["Great summary of the thread.", "api, design, testing"])
|
||||
|
||||
key = (400, 99)
|
||||
_thread_history[key] = [
|
||||
{"role": "user", "content": "question one"},
|
||||
{"role": "assistant", "content": "answer one"},
|
||||
{"role": "user", "content": "question two"},
|
||||
{"role": "assistant", "content": "answer two"},
|
||||
]
|
||||
|
||||
update = _make_update(chat_id=400, thread_id=99)
|
||||
ctx = _make_context(settings, llm=mock_llm, store=store)
|
||||
await flush_handler(update, ctx)
|
||||
|
||||
# LLM called twice: once for summary, once for tags
|
||||
assert mock_llm.chat.await_count == 2
|
||||
# Summary persisted
|
||||
store.save.assert_called_once()
|
||||
saved: ThreadSummary = store.save.call_args.args[0]
|
||||
assert saved.chat_id == 400
|
||||
assert saved.thread_id == 99
|
||||
assert saved.summary == "Great summary of the thread."
|
||||
assert saved.message_count == 2 # 2 user turns
|
||||
assert saved.tags == ["api", "design", "testing"]
|
||||
|
||||
# In-memory history replaced with compressed context
|
||||
assert len(_thread_history[key]) == 1
|
||||
assert _thread_history[key][0]["role"] == "system"
|
||||
assert "Great summary" in _thread_history[key][0]["content"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# /recall handler tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_recall_in_thread_returns_stored_summary(tmp_path):
|
||||
"""recall_handler inside a thread returns the stored summary for that thread."""
|
||||
settings = _make_settings()
|
||||
store = ThreadMemoryStore(tmp_path / "mem.json")
|
||||
store.save(ThreadSummary(chat_id=500, thread_id=11, summary="recap text", message_count=3))
|
||||
|
||||
update = _make_update(chat_id=500, thread_id=11)
|
||||
ctx = _make_context(settings, store=store)
|
||||
await recall_handler(update, ctx)
|
||||
|
||||
update.message.reply_text.assert_awaited_once()
|
||||
text = update.message.reply_text.call_args.args[0]
|
||||
assert "recap text" in text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_recall_in_thread_no_summary_guides_user(tmp_path):
|
||||
"""recall_handler inside a thread with no summary tells user to /flush first."""
|
||||
settings = _make_settings()
|
||||
store = ThreadMemoryStore(tmp_path / "mem.json")
|
||||
|
||||
update = _make_update(chat_id=600, thread_id=22)
|
||||
ctx = _make_context(settings, store=store)
|
||||
await recall_handler(update, ctx)
|
||||
|
||||
update.message.reply_text.assert_awaited_once()
|
||||
text = update.message.reply_text.call_args.args[0]
|
||||
assert "flush" in text.lower()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_recall_outside_thread_lists_all(tmp_path):
|
||||
"""recall_handler outside a thread lists all stored summaries."""
|
||||
settings = _make_settings()
|
||||
store = ThreadMemoryStore(tmp_path / "mem.json")
|
||||
store.save(ThreadSummary(chat_id=1, thread_id=1, summary="alpha", message_count=1))
|
||||
store.save(ThreadSummary(chat_id=1, thread_id=2, summary="beta", message_count=2))
|
||||
|
||||
update = _make_update(thread_id=None) # no thread
|
||||
ctx = _make_context(settings, store=store)
|
||||
await recall_handler(update, ctx)
|
||||
|
||||
update.message.reply_text.assert_awaited_once()
|
||||
text = update.message.reply_text.call_args.args[0]
|
||||
assert "alpha" in text or "beta" in text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_recall_outside_thread_empty_store(tmp_path):
|
||||
"""recall_handler outside a thread with no summaries guides user."""
|
||||
settings = _make_settings()
|
||||
store = ThreadMemoryStore(tmp_path / "mem.json")
|
||||
|
||||
update = _make_update(thread_id=None)
|
||||
ctx = _make_context(settings, store=store)
|
||||
await recall_handler(update, ctx)
|
||||
|
||||
update.message.reply_text.assert_awaited_once()
|
||||
text = update.message.reply_text.call_args.args[0]
|
||||
assert "flush" in text.lower()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_recall_with_query_returns_matching_summaries(tmp_path):
|
||||
"""recall_handler with a query argument searches the knowledge base by keyword."""
|
||||
settings = _make_settings()
|
||||
store = ThreadMemoryStore(tmp_path / "mem.json")
|
||||
store.save(ThreadSummary(
|
||||
chat_id=1, thread_id=1, summary="API authentication discussion",
|
||||
message_count=2, tags=["api", "auth"],
|
||||
))
|
||||
store.save(ThreadSummary(
|
||||
chat_id=1, thread_id=2, summary="Database schema planning",
|
||||
message_count=3, tags=["database", "schema"],
|
||||
))
|
||||
|
||||
update = _make_update(thread_id=None)
|
||||
ctx = _make_context(settings, store=store)
|
||||
ctx.args = ["api"]
|
||||
await recall_handler(update, ctx)
|
||||
|
||||
update.message.reply_text.assert_awaited_once()
|
||||
text = update.message.reply_text.call_args.args[0]
|
||||
assert "API authentication" in text
|
||||
assert "Database" not in text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_recall_with_query_no_match(tmp_path):
|
||||
"""recall_handler with a query that matches nothing tells the user."""
|
||||
settings = _make_settings()
|
||||
store = ThreadMemoryStore(tmp_path / "mem.json")
|
||||
store.save(ThreadSummary(
|
||||
chat_id=1, thread_id=1, summary="recap", message_count=1, tags=["database"],
|
||||
))
|
||||
|
||||
update = _make_update(thread_id=None)
|
||||
ctx = _make_context(settings, store=store)
|
||||
ctx.args = ["deployment"]
|
||||
await recall_handler(update, ctx)
|
||||
|
||||
update.message.reply_text.assert_awaited_once()
|
||||
text = update.message.reply_text.call_args.args[0]
|
||||
assert "deployment" in text.lower() or "no memories" in text.lower()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_message_handler_injects_kb_context_transiently(tmp_path):
|
||||
"""message_handler injects relevant KB summaries as transient context without storing them."""
|
||||
settings = _make_settings()
|
||||
mock_llm = MagicMock(spec=LLMClient)
|
||||
mock_llm.chat = AsyncMock(return_value="reply with context")
|
||||
|
||||
store = ThreadMemoryStore(tmp_path / "mem.json")
|
||||
store.save(ThreadSummary(
|
||||
chat_id=1, thread_id=1, summary="Previous API discussion",
|
||||
message_count=2, tags=["api"],
|
||||
))
|
||||
|
||||
user_id = 999
|
||||
from steward.bot.telegram import _history
|
||||
_history[user_id].clear()
|
||||
|
||||
update = _make_update(user_id=user_id, text="tell me about the api work")
|
||||
ctx = _make_context(settings, llm=mock_llm, store=store)
|
||||
await message_handler(update, ctx)
|
||||
|
||||
# LLM was called
|
||||
mock_llm.chat.assert_awaited_once()
|
||||
call_kwargs = mock_llm.chat.call_args
|
||||
# The history passed to the LLM should contain the KB context message
|
||||
history_arg = call_kwargs.kwargs.get("history") or (
|
||||
call_kwargs.args[1] if len(call_kwargs.args) > 1 else None
|
||||
)
|
||||
assert history_arg is not None
|
||||
assert any("knowledge base" in m.get("content", "").lower() for m in history_arg)
|
||||
|
||||
# But the KB message must NOT be stored in _history
|
||||
assert all(
|
||||
"knowledge base" not in m.get("content", "").lower()
|
||||
for m in _history[user_id]
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_message_handler_no_kb_injection_when_no_match(tmp_path):
|
||||
"""message_handler does not inject KB context when no summaries match."""
|
||||
settings = _make_settings()
|
||||
mock_llm = MagicMock(spec=LLMClient)
|
||||
mock_llm.chat = AsyncMock(return_value="plain reply")
|
||||
|
||||
store = ThreadMemoryStore(tmp_path / "mem.json")
|
||||
# Store a summary with unrelated tags
|
||||
store.save(ThreadSummary(
|
||||
chat_id=1, thread_id=1, summary="Database recap", message_count=1, tags=["database"],
|
||||
))
|
||||
|
||||
user_id = 888
|
||||
from steward.bot.telegram import _history
|
||||
_history[user_id].clear()
|
||||
|
||||
update = _make_update(user_id=user_id, text="what is the weather like?")
|
||||
ctx = _make_context(settings, llm=mock_llm, store=store)
|
||||
await message_handler(update, ctx)
|
||||
|
||||
call_kwargs = mock_llm.chat.call_args
|
||||
history_arg = call_kwargs.kwargs.get("history") or []
|
||||
# No KB context injected
|
||||
assert not any("knowledge base" in m.get("content", "").lower() for m in history_arg)
|
||||
499
tests/test_tools.py
Normal file
499
tests/test_tools.py
Normal file
@ -0,0 +1,499 @@
|
||||
"""Tests for steward.tools.client and LLM tool-calling integration."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
import respx
|
||||
from httpx import Response
|
||||
|
||||
from steward.config import Settings
|
||||
from steward.llm.client import LLMClient
|
||||
from steward.tools.client import (
|
||||
ToolClient,
|
||||
_operation_id,
|
||||
_resolve_operation,
|
||||
_schema_to_json_schema,
|
||||
_spec_to_openai_tools,
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fixtures
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
SIMPLE_SPEC: dict[str, Any] = {
|
||||
"openapi": "3.0.0",
|
||||
"info": {"title": "Test", "version": "1.0.0"},
|
||||
"paths": {
|
||||
"/items": {
|
||||
"get": {
|
||||
"operationId": "list_items",
|
||||
"summary": "List all items",
|
||||
"parameters": [
|
||||
{
|
||||
"name": "limit",
|
||||
"in": "query",
|
||||
"required": False,
|
||||
"schema": {"type": "integer"},
|
||||
"description": "Max results",
|
||||
}
|
||||
],
|
||||
"responses": {"200": {"description": "OK"}},
|
||||
}
|
||||
},
|
||||
"/items/{item_id}": {
|
||||
"get": {
|
||||
"operationId": "get_item",
|
||||
"summary": "Get a single item",
|
||||
"parameters": [
|
||||
{
|
||||
"name": "item_id",
|
||||
"in": "path",
|
||||
"required": True,
|
||||
"schema": {"type": "string"},
|
||||
}
|
||||
],
|
||||
"responses": {"200": {"description": "OK"}},
|
||||
}
|
||||
},
|
||||
"/items/create": {
|
||||
"post": {
|
||||
"operationId": "create_item",
|
||||
"summary": "Create an item",
|
||||
"requestBody": {
|
||||
"required": True,
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"name": {"type": "string", "description": "Item name"},
|
||||
"value": {"type": "integer"},
|
||||
},
|
||||
"required": ["name"],
|
||||
}
|
||||
}
|
||||
},
|
||||
},
|
||||
"responses": {"201": {"description": "Created"}},
|
||||
}
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _make_llm_settings(**kwargs: Any) -> Settings:
|
||||
defaults = dict(
|
||||
telegram_bot_token="t",
|
||||
openai_api_key="k",
|
||||
openai_model="gpt-4o",
|
||||
openai_system_prompt="You are Steward.",
|
||||
)
|
||||
defaults.update(kwargs)
|
||||
return Settings(**defaults)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _operation_id
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_generate_operation_id_simple():
|
||||
op: dict[str, Any] = {}
|
||||
assert _operation_id("get", "/items", op) == "get_items"
|
||||
|
||||
|
||||
def test_generate_operation_id_with_path_param():
|
||||
op: dict[str, Any] = {}
|
||||
assert _operation_id("get", "/items/{item_id}", op) == "get_items_item_id"
|
||||
|
||||
|
||||
def test_generate_operation_id_root():
|
||||
op: dict[str, Any] = {}
|
||||
assert _operation_id("get", "/", op) == "get"
|
||||
|
||||
|
||||
def test_operation_id_uses_declared():
|
||||
op = {"operationId": "my_op"}
|
||||
assert _operation_id("get", "/foo", op) == "my_op"
|
||||
|
||||
|
||||
def test_operation_id_generates_from_path():
|
||||
op: dict[str, Any] = {}
|
||||
assert _operation_id("post", "/foo/bar", op) == "post_foo_bar"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _schema_to_json_schema
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_schema_to_json_schema_basic():
|
||||
schema = {"type": "string", "description": "A name"}
|
||||
result = _schema_to_json_schema(schema)
|
||||
assert result == {"type": "string", "description": "A name"}
|
||||
|
||||
|
||||
def test_schema_to_json_schema_resolves_ref():
|
||||
components = {"schemas": {"Foo": {"type": "integer"}}}
|
||||
schema = {"$ref": "#/components/schemas/Foo"}
|
||||
result = _schema_to_json_schema(schema, components)
|
||||
assert result == {"type": "integer"}
|
||||
|
||||
|
||||
def test_schema_to_json_schema_nested_properties():
|
||||
schema = {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"name": {"type": "string"},
|
||||
"count": {"type": "integer"},
|
||||
},
|
||||
"required": ["name"],
|
||||
}
|
||||
result = _schema_to_json_schema(schema)
|
||||
assert result["type"] == "object"
|
||||
assert "name" in result["properties"]
|
||||
assert result["required"] == ["name"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _spec_to_openai_tools
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_spec_to_openai_tools_count():
|
||||
tools = _spec_to_openai_tools(SIMPLE_SPEC)
|
||||
assert len(tools) == 3
|
||||
|
||||
|
||||
def test_spec_to_openai_tools_structure():
|
||||
tools = _spec_to_openai_tools(SIMPLE_SPEC)
|
||||
tool = next(t for t in tools if t["function"]["name"] == "list_items")
|
||||
assert tool["type"] == "function"
|
||||
assert tool["function"]["description"] == "List all items"
|
||||
params = tool["function"]["parameters"]
|
||||
assert params["type"] == "object"
|
||||
assert "limit" in params["properties"]
|
||||
|
||||
|
||||
def test_spec_to_openai_tools_path_param_required():
|
||||
tools = _spec_to_openai_tools(SIMPLE_SPEC)
|
||||
tool = next(t for t in tools if t["function"]["name"] == "get_item")
|
||||
assert "item_id" in tool["function"]["parameters"]["properties"]
|
||||
assert "item_id" in tool["function"]["parameters"]["required"]
|
||||
|
||||
|
||||
def test_spec_to_openai_tools_request_body():
|
||||
tools = _spec_to_openai_tools(SIMPLE_SPEC)
|
||||
tool = next(t for t in tools if t["function"]["name"] == "create_item")
|
||||
props = tool["function"]["parameters"]["properties"]
|
||||
assert "name" in props
|
||||
assert "value" in props
|
||||
|
||||
|
||||
def test_spec_to_openai_tools_empty_spec():
|
||||
assert _spec_to_openai_tools({}) == []
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _resolve_operation
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_resolve_operation_query_params():
|
||||
path, method, path_p, query_p, body = _resolve_operation(
|
||||
SIMPLE_SPEC, "list_items", {"limit": 5}
|
||||
)
|
||||
assert path == "/items"
|
||||
assert method == "get"
|
||||
assert path_p == {}
|
||||
assert query_p == {"limit": 5}
|
||||
assert body == {}
|
||||
|
||||
|
||||
def test_resolve_operation_path_params():
|
||||
path, method, path_p, query_p, body = _resolve_operation(
|
||||
SIMPLE_SPEC, "get_item", {"item_id": "abc"}
|
||||
)
|
||||
assert path == "/items/{item_id}"
|
||||
assert method == "get"
|
||||
assert path_p == {"item_id": "abc"}
|
||||
assert query_p == {}
|
||||
assert body == {}
|
||||
|
||||
|
||||
def test_resolve_operation_body():
|
||||
path, method, path_p, query_p, body = _resolve_operation(
|
||||
SIMPLE_SPEC, "create_item", {"name": "foo", "value": 42}
|
||||
)
|
||||
assert path == "/items/create"
|
||||
assert method == "post"
|
||||
assert path_p == {}
|
||||
assert query_p == {}
|
||||
assert body == {"name": "foo", "value": 42}
|
||||
|
||||
|
||||
def test_resolve_operation_not_found():
|
||||
with pytest.raises(ValueError, match="not found"):
|
||||
_resolve_operation(SIMPLE_SPEC, "nonexistent_op", {})
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ToolClient
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_client_get_spec():
|
||||
"""get_spec() fetches from /openapi.json and caches the result."""
|
||||
with respx.mock:
|
||||
respx.get("http://tools.local/openapi.json").mock(
|
||||
return_value=Response(200, json=SIMPLE_SPEC)
|
||||
)
|
||||
client = ToolClient("http://tools.local")
|
||||
spec = await client.get_spec()
|
||||
assert spec["openapi"] == "3.0.0"
|
||||
# Second call uses cache – no additional HTTP request
|
||||
spec2 = await client.get_spec()
|
||||
assert spec2 is spec
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_client_get_tools():
|
||||
with respx.mock:
|
||||
respx.get("http://tools.local/openapi.json").mock(
|
||||
return_value=Response(200, json=SIMPLE_SPEC)
|
||||
)
|
||||
client = ToolClient("http://tools.local")
|
||||
tools = await client.get_tools()
|
||||
assert len(tools) == 3
|
||||
# Cached
|
||||
tools2 = await client.get_tools()
|
||||
assert tools2 is tools
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_client_call_get_with_query():
|
||||
with respx.mock:
|
||||
respx.get("http://tools.local/openapi.json").mock(
|
||||
return_value=Response(200, json=SIMPLE_SPEC)
|
||||
)
|
||||
respx.get("http://tools.local/items").mock(
|
||||
return_value=Response(200, json=[{"id": 1}])
|
||||
)
|
||||
client = ToolClient("http://tools.local")
|
||||
result = await client.call("list_items", {"limit": 10})
|
||||
assert "id" in result
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_client_call_path_substitution():
|
||||
with respx.mock:
|
||||
respx.get("http://tools.local/openapi.json").mock(
|
||||
return_value=Response(200, json=SIMPLE_SPEC)
|
||||
)
|
||||
respx.get("http://tools.local/items/abc123").mock(
|
||||
return_value=Response(200, json={"id": "abc123"})
|
||||
)
|
||||
client = ToolClient("http://tools.local")
|
||||
result = await client.call("get_item", {"item_id": "abc123"})
|
||||
assert "abc123" in result
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_client_includes_auth_header():
|
||||
"""ToolClient sends Authorization header when api_key is provided."""
|
||||
captured_headers: dict[str, str] = {}
|
||||
|
||||
def capture(request: Any, *args: Any, **kwargs: Any) -> Response:
|
||||
captured_headers.update(dict(request.headers))
|
||||
return Response(200, json=SIMPLE_SPEC)
|
||||
|
||||
with respx.mock:
|
||||
respx.get("http://tools.local/openapi.json").mock(side_effect=capture)
|
||||
client = ToolClient("http://tools.local", api_key="secret-token")
|
||||
await client.get_spec()
|
||||
|
||||
assert "authorization" in {k.lower() for k in captured_headers}
|
||||
auth = next(v for k, v in captured_headers.items() if k.lower() == "authorization")
|
||||
assert auth.startswith("Bearer ")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_client_returns_plain_text_for_non_json():
|
||||
with respx.mock:
|
||||
respx.get("http://tools.local/openapi.json").mock(
|
||||
return_value=Response(200, json=SIMPLE_SPEC)
|
||||
)
|
||||
respx.get("http://tools.local/items").mock(
|
||||
return_value=Response(200, text="plain result", headers={"content-type": "text/plain"})
|
||||
)
|
||||
client = ToolClient("http://tools.local")
|
||||
result = await client.call("list_items", {})
|
||||
assert result == "plain result"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# LLMClient.chat_with_tools
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_tool_response(tool_name: str, tool_call_id: str, args: str) -> MagicMock:
|
||||
"""Create a mock OpenAI response that requests a tool call."""
|
||||
tc = MagicMock()
|
||||
tc.id = tool_call_id
|
||||
tc.function.name = tool_name
|
||||
tc.function.arguments = args
|
||||
|
||||
msg = MagicMock()
|
||||
msg.content = None
|
||||
msg.tool_calls = [tc]
|
||||
|
||||
choice = MagicMock()
|
||||
choice.message = msg
|
||||
|
||||
resp = MagicMock()
|
||||
resp.choices = [choice]
|
||||
return resp
|
||||
|
||||
|
||||
def _make_text_response(content: str) -> MagicMock:
|
||||
"""Create a mock OpenAI response with plain text content."""
|
||||
msg = MagicMock()
|
||||
msg.content = content
|
||||
msg.tool_calls = None
|
||||
|
||||
choice = MagicMock()
|
||||
choice.message = msg
|
||||
|
||||
resp = MagicMock()
|
||||
resp.choices = [choice]
|
||||
return resp
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_with_tools_no_tool_calls():
|
||||
"""When the LLM returns no tool calls, reply is returned directly."""
|
||||
settings = _make_llm_settings()
|
||||
llm = LLMClient(settings)
|
||||
|
||||
mock_tool_client = AsyncMock()
|
||||
mock_tool_client.get_tools = AsyncMock(return_value=[])
|
||||
|
||||
with patch.object(
|
||||
llm._client.chat.completions,
|
||||
"create",
|
||||
new_callable=AsyncMock,
|
||||
return_value=_make_text_response("Hello world"),
|
||||
):
|
||||
result = await llm.chat_with_tools("Hi", mock_tool_client)
|
||||
|
||||
assert result == "Hello world"
|
||||
mock_tool_client.call.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_with_tools_executes_tool_call():
|
||||
"""Tool call is executed and result fed back; final text is returned."""
|
||||
settings = _make_llm_settings()
|
||||
llm = LLMClient(settings)
|
||||
|
||||
tool_response = _make_tool_response("list_items", "call-1", json.dumps({"limit": 5}))
|
||||
final_response = _make_text_response("Here are the items: [foo, bar]")
|
||||
|
||||
mock_tool_client = AsyncMock()
|
||||
mock_tool_client.get_tools = AsyncMock(return_value=[{"type": "function"}])
|
||||
mock_tool_client.call = AsyncMock(return_value='[{"name": "foo"}, {"name": "bar"}]')
|
||||
|
||||
with patch.object(
|
||||
llm._client.chat.completions,
|
||||
"create",
|
||||
new_callable=AsyncMock,
|
||||
side_effect=[tool_response, final_response],
|
||||
):
|
||||
result = await llm.chat_with_tools("List items", mock_tool_client)
|
||||
|
||||
assert result == "Here are the items: [foo, bar]"
|
||||
mock_tool_client.call.assert_awaited_once_with("list_items", {"limit": 5})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_with_tools_handles_tool_error():
|
||||
"""When tool execution fails the error is fed back to the LLM gracefully."""
|
||||
settings = _make_llm_settings()
|
||||
llm = LLMClient(settings)
|
||||
|
||||
tool_response = _make_tool_response("list_items", "call-1", json.dumps({}))
|
||||
final_response = _make_text_response("I could not retrieve items due to an error.")
|
||||
|
||||
mock_tool_client = AsyncMock()
|
||||
mock_tool_client.get_tools = AsyncMock(return_value=[{"type": "function"}])
|
||||
mock_tool_client.call = AsyncMock(side_effect=RuntimeError("Server unreachable"))
|
||||
|
||||
with patch.object(
|
||||
llm._client.chat.completions,
|
||||
"create",
|
||||
new_callable=AsyncMock,
|
||||
side_effect=[tool_response, final_response],
|
||||
) as mock_create:
|
||||
result = await llm.chat_with_tools("List items", mock_tool_client)
|
||||
|
||||
assert "error" in result.lower() or "could not" in result.lower()
|
||||
# The second LLM call should include the tool error in messages
|
||||
second_call_messages = mock_create.call_args_list[1].kwargs["messages"]
|
||||
tool_result_msg = next(m for m in second_call_messages if m.get("role") == "tool")
|
||||
assert "Tool call failed" in tool_result_msg["content"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_with_tools_respects_max_rounds():
|
||||
"""chat_with_tools exits after max_tool_rounds even if LLM keeps calling tools."""
|
||||
settings = _make_llm_settings()
|
||||
llm = LLMClient(settings)
|
||||
|
||||
tool_response = _make_tool_response("list_items", "call-1", json.dumps({}))
|
||||
|
||||
mock_tool_client = AsyncMock()
|
||||
mock_tool_client.get_tools = AsyncMock(return_value=[{"type": "function"}])
|
||||
mock_tool_client.call = AsyncMock(return_value="[]")
|
||||
|
||||
# Always returns a tool call – never a final answer
|
||||
with patch.object(
|
||||
llm._client.chat.completions,
|
||||
"create",
|
||||
new_callable=AsyncMock,
|
||||
return_value=tool_response,
|
||||
) as mock_create:
|
||||
await llm.chat_with_tools("List items", mock_tool_client, max_tool_rounds=3)
|
||||
|
||||
assert mock_create.call_count == 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_with_tools_passes_history():
|
||||
"""History is included in the first LLM request."""
|
||||
settings = _make_llm_settings()
|
||||
llm = LLMClient(settings)
|
||||
|
||||
mock_tool_client = AsyncMock()
|
||||
mock_tool_client.get_tools = AsyncMock(return_value=[])
|
||||
|
||||
history = [
|
||||
{"role": "user", "content": "previous"},
|
||||
{"role": "assistant", "content": "answer"},
|
||||
]
|
||||
|
||||
with patch.object(
|
||||
llm._client.chat.completions,
|
||||
"create",
|
||||
new_callable=AsyncMock,
|
||||
return_value=_make_text_response("reply"),
|
||||
) as mock_create:
|
||||
await llm.chat_with_tools("new question", mock_tool_client, history=history)
|
||||
|
||||
messages = mock_create.call_args.kwargs["messages"]
|
||||
roles = [m["role"] for m in messages]
|
||||
assert roles == ["system", "user", "assistant", "user"]
|
||||
Loading…
x
Reference in New Issue
Block a user