Compare commits
@@ -0,0 +1 @@
|
||||
{"sessionId":"45d44ada-86af-4207-b3dd-de510e521157","pid":20439,"acquiredAt":1773554855575}
|
||||
@@ -8,6 +8,7 @@ on:
|
||||
permissions:
|
||||
contents: write
|
||||
packages: write
|
||||
id-token: write
|
||||
|
||||
env:
|
||||
GO_VERSION: "1.25"
|
||||
@@ -217,3 +218,28 @@ jobs:
|
||||
labels: ${{ steps.meta.outputs.labels }}
|
||||
cache-from: type=gha
|
||||
cache-to: type=gha,mode=max
|
||||
|
||||
mcp-registry:
|
||||
name: Publish to MCP Registry
|
||||
needs: release
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Extract version from tag
|
||||
id: version
|
||||
run: echo "VERSION=${GITHUB_REF_NAME#v}" >> "$GITHUB_OUTPUT"
|
||||
|
||||
- name: Install mcp-publisher
|
||||
run: |
|
||||
curl -L "https://github.com/modelcontextprotocol/registry/releases/latest/download/mcp-publisher_linux_amd64.tar.gz" | tar xz mcp-publisher
|
||||
|
||||
- name: Authenticate to MCP Registry
|
||||
run: ./mcp-publisher login github-oidc
|
||||
|
||||
- name: Update version in server.json
|
||||
run: |
|
||||
jq --arg v "${{ steps.version.outputs.VERSION }}" '.version = $v' server.json > server.tmp && mv server.tmp server.json
|
||||
|
||||
- name: Publish to MCP Registry
|
||||
run: ./mcp-publisher publish
|
||||
|
||||
@@ -100,6 +100,12 @@ make lint # Run linters
|
||||
- SQLite (modernc.org/sqlite, pure Go) — single DB file in `--data` directory (004-embeddings-retention-inbox)
|
||||
- Go 1.25+ (per go.mod) + spf13/cobra (CLI), go-chi/chi (HTTP), mark3labs/mcp-go (MCP) (006-admin-cli-docker-fixes)
|
||||
- modernc.org/sqlite (pure Go, zero CGO) (006-admin-cli-docker-fixes)
|
||||
- Go 1.25+ (per go.mod) + go-chi/chi (HTTP), mark3labs/mcp-go (MCP), ory/fosite (OAuth), spf13/cobra (CLI), modernc.org/sqlite (storage), TFMV/hnsw (vectors). NEW: coreos/go-oidc/v3 (OIDC), golang.org/x/oauth2 (OAuth client) (007-platform-features-bundle)
|
||||
- Go 1.25+ (backend), SvelteKit 2 + Svelte 5 (frontend), SvelteKit (website) + go-chi/chi (HTTP), mark3labs/mcp-go (MCP), modernc.org/sqlite (storage), SherClockHolmes/webpush-go (push notifications — NEW) (008-webui-pwa-analytics)
|
||||
- SQLite (existing DB, 1 new migration for push_subscriptions), localStorage (font size) (008-webui-pwa-analytics)
|
||||
- Go 1.25+ (backend), Svelte 5 + Tailwind (frontend) + go-chi/chi (HTTP), mark3labs/mcp-go (MCP), modernc.org/sqlite (storage), spf13/cobra (CLI) (009-attachments-threads)
|
||||
- SQLite (modernc.org/sqlite, pure Go) + content-addressable filesystem (SHA-256) (009-attachments-threads)
|
||||
- SQLite (modernc.org/sqlite, pure Go) — new migration 013_reactions.sql (010-reactions-workflows)
|
||||
|
||||
## Recent Changes
|
||||
- 002-mcp-auth-ux-polish: Added Go 1.23+ + ory/fosite (OAuth 2.1), mark3labs/mcp-go (MCP server), go-chi/chi (HTTP), Svelte 5 + Tailwind (Web UI)
|
||||
|
||||
+1
-1
@@ -19,7 +19,7 @@ RUN CGO_ENABLED=0 GOOS=linux go build -ldflags="-s -w -X main.version=${VERSION}
|
||||
|
||||
# Stage 3: Runtime
|
||||
FROM alpine:3.19
|
||||
RUN apk add --no-cache ca-certificates tzdata
|
||||
RUN apk add --no-cache ca-certificates tzdata && touch /.dockerenv
|
||||
COPY --from=go-builder /synapbus /synapbus
|
||||
EXPOSE 8080
|
||||
VOLUME ["/data"]
|
||||
|
||||
+80
-75
@@ -1,96 +1,101 @@
|
||||
# Autonomous Implementation Summary
|
||||
# Autonomous Implementation Summary: Message Reactions & Workflow States
|
||||
|
||||
**Feature**: Admin CLI & Docker Fixes
|
||||
**Branch**: `006-admin-cli-docker-fixes`
|
||||
**Date**: 2026-03-15
|
||||
**Status**: COMPLETE — 8 of 8 tasks implemented, all tests pass, binary builds
|
||||
**Branch**: `010-reactions-workflows`
|
||||
**Date**: 2026-03-18
|
||||
**Status**: Complete (StalemateWorker extension deferred)
|
||||
|
||||
## What Was Built
|
||||
|
||||
### 1. Alpine Docker Base Image (T06)
|
||||
### Message Reactions
|
||||
- **Toggle semantics**: Add a reaction → added. Add same reaction again → removed. One per type per agent per message.
|
||||
- **5 reaction types**: approve, reject, in_progress, done, published
|
||||
- **Metadata support**: JSON metadata on reactions (e.g., `{"url": "https://..."}` for published)
|
||||
- **100-reaction limit** per message (safety)
|
||||
|
||||
**Problem**: `scratch` base image has no shell — `kubectl exec` into the pod can't run admin CLI commands.
|
||||
### Workflow State Derivation
|
||||
- State computed from reactions: published > done > rejected > in_progress > approved > proposed
|
||||
- No denormalization — state derived on read from reaction list
|
||||
- Channel messages with no reactions → "proposed" state
|
||||
- Terminal states (rejected, done, published) don't trigger stalemate checks
|
||||
|
||||
**Solution**: Changed `FROM scratch` to `FROM alpine:3.19` in the runtime stage. Alpine provides `/bin/sh` and a working process environment. TLS certs and timezone data are now installed via `apk` instead of copied from the builder stage.
|
||||
### Channel Workflow Settings
|
||||
- `auto_approve` — skip proposed state for new messages
|
||||
- `stalemate_remind_after` — duration before reminder DM (default 24h)
|
||||
- `stalemate_escalate_after` — duration before escalation to #approvals (default 72h)
|
||||
|
||||
**File Modified**: `Dockerfile`
|
||||
### REST API
|
||||
- `POST /api/messages/{id}/reactions` — toggle reaction (add or remove)
|
||||
- `GET /api/messages/{id}/reactions` — get reactions + workflow state
|
||||
- `DELETE /api/messages/{id}/reactions/{reaction}` — remove reaction
|
||||
- `PUT /api/channels/{name}/settings` — update workflow settings
|
||||
- `GET /api/channels/{name}/messages/by-state?state=X` — list messages by state
|
||||
|
||||
### 2. `synapbus channels create` CLI Command (T02, T04)
|
||||
### MCP Tools (via execute bridge)
|
||||
- `react` — add/toggle reaction on a message
|
||||
- `unreact` — remove a reaction
|
||||
- `get_reactions` — query reactions and workflow state
|
||||
- `list_by_state` — list messages by workflow state in a channel
|
||||
|
||||
**Problem**: No CLI command to create channels — had to use REST API with session cookies.
|
||||
### Web UI
|
||||
- **WorkflowBadge** component: colored pills (yellow/green/blue/red/gray/cyan) per state
|
||||
- **ReactionPills** component: grouped reaction pills with count, agent names on hover, click-to-toggle
|
||||
- Published reactions with URL show clickable link icon
|
||||
- Integrated into channel message view
|
||||
|
||||
**Solution**: Added `channels.create` admin socket handler and `synapbus channels create` cobra command:
|
||||
- `--name` (required): Channel name
|
||||
- `--description` (optional): Channel description
|
||||
- Creates channel via the channel service with `created_by: "system"`, type `"standard"`
|
||||
- Returns channel details as JSON
|
||||
### Admin CLI
|
||||
- `synapbus channels update --name X --auto-approve=true --stalemate-remind-after=12h --stalemate-escalate-after=48h`
|
||||
|
||||
**Files Modified**: `cmd/synapbus/admin.go`, `internal/admin/socket.go`
|
||||
## Files Created/Modified
|
||||
|
||||
### 3. `synapbus channels join` CLI Command (T03, T05)
|
||||
|
||||
**Problem**: No CLI command to add agents to channels.
|
||||
|
||||
**Solution**: Added `channels.join` admin socket handler and `synapbus channels join` cobra command:
|
||||
- `--channel` (required): Channel name to join
|
||||
- `--agent` (required): Agent name to add
|
||||
- Looks up channel by name, calls `JoinChannel` (idempotent)
|
||||
- Reports `"joined"` or `"already_member"` status
|
||||
|
||||
**Files Modified**: `cmd/synapbus/admin.go`, `internal/admin/socket.go`
|
||||
|
||||
### 4. Absolute Default Socket Path (T01)
|
||||
|
||||
**Problem**: Default `./data/synapbus.sock` is confusing in containers where CWD varies.
|
||||
|
||||
**Solution**: Changed default socket path from `./data/synapbus.sock` to `/data/synapbus.sock` in both the `--socket` flag definition and the `SYNAPBUS_SOCKET` env var comparison.
|
||||
|
||||
**File Modified**: `cmd/synapbus/admin.go`
|
||||
|
||||
## Tests Added (T07)
|
||||
|
||||
| Test | Description |
|
||||
### New Files
|
||||
| File | Description |
|
||||
|------|-------------|
|
||||
| `TestChannelsCreateCommandRegistered` | Verifies `channels create` subcommand exists |
|
||||
| `TestChannelsCreateRequiredFlags` | Verifies `--name` is required, `--description` is optional |
|
||||
| `TestChannelsJoinCommandRegistered` | Verifies `channels join` subcommand exists |
|
||||
| `TestChannelsJoinRequiredFlags` | Verifies `--channel` and `--agent` are both required |
|
||||
| `TestDefaultSocketPath` | Verifies default is `/data/synapbus.sock` |
|
||||
|
||||
## Verification Results (T08)
|
||||
|
||||
| Check | Result |
|
||||
|-------|--------|
|
||||
| `go build ./...` | PASS |
|
||||
| `go test ./...` | ALL PASS (24 packages, 0 failures) |
|
||||
| Zero CGO | Confirmed (CGO_ENABLED=0 in Dockerfile) |
|
||||
| No regressions | All 14 existing CLI tests still pass |
|
||||
|
||||
## Files Changed
|
||||
| `internal/storage/schema/013_reactions.sql` | Migration: message_reactions table + channel columns |
|
||||
| `internal/reactions/model.go` | Reaction types, state derivation, constants |
|
||||
| `internal/reactions/store.go` | SQLite CRUD for reactions |
|
||||
| `internal/reactions/service.go` | Business logic: toggle, remove, get, list by state |
|
||||
| `internal/reactions/model_test.go` | 23 test cases for model functions |
|
||||
| `internal/reactions/store_test.go` | 6 test functions for store operations |
|
||||
| `internal/api/reactions_handler.go` | REST API handlers for reactions |
|
||||
| `web/src/lib/components/WorkflowBadge.svelte` | Colored state badge component |
|
||||
| `web/src/lib/components/ReactionPills.svelte` | Reaction toggle pills component |
|
||||
|
||||
### Modified Files
|
||||
| File | Changes |
|
||||
|------|---------|
|
||||
| `Dockerfile` | `FROM scratch` → `FROM alpine:3.19` + `apk add --no-cache ca-certificates tzdata` |
|
||||
| `cmd/synapbus/admin.go` | Default socket `/data/synapbus.sock`, `channels create` + `channels join` commands |
|
||||
| `cmd/synapbus/admin_test.go` | 5 new tests for commands, flags, and default socket |
|
||||
| `internal/admin/socket.go` | `channels.create` + `channels.join` handlers, `channels` import |
|
||||
| `internal/messaging/types.go` | Added WorkflowState, Reactions, ReactionInfo to Message |
|
||||
| `internal/messaging/service.go` | Added ReactionEnricher interface, enrichment in EnrichMessages |
|
||||
| `internal/channels/types.go` | Added AutoApprove, StalemateRemindAfter, StalemateEscalateAfter, ChannelSettings |
|
||||
| `internal/channels/store.go` | Updated SELECT queries for new columns, added UpdateChannelSettings |
|
||||
| `internal/channels/service.go` | Added UpdateChannelSettings method |
|
||||
| `internal/api/router.go` | Registered reaction and channel settings routes |
|
||||
| `internal/api/channels_handler.go` | Added UpdateSettings, ListByState handlers |
|
||||
| `internal/mcp/bridge.go` | Added react/unreact/get_reactions/list_by_state bridge methods |
|
||||
| `internal/mcp/tools_hybrid.go` | Added reactionService to registrar |
|
||||
| `internal/mcp/server.go` | Added reactionService parameter |
|
||||
| `internal/actions/registry.go` | Registered 4 new reaction actions |
|
||||
| `cmd/synapbus/main.go` | Wired reaction service, adapter, passed to router+MCP |
|
||||
| `cmd/synapbus/admin.go` | Added channels update CLI command |
|
||||
| `internal/admin/socket.go` | Added channels.update_settings handler |
|
||||
| `web/src/lib/api/client.ts` | Added reactions.toggle/get methods |
|
||||
| `web/src/routes/channels/[name]/+page.svelte` | Integrated WorkflowBadge + ReactionPills |
|
||||
|
||||
## CLI Commands Added
|
||||
## Test Results
|
||||
|
||||
| Command | Description |
|
||||
|---------|-------------|
|
||||
| `synapbus channels create --name X [--description Y]` | Create a new channel |
|
||||
| `synapbus channels join --channel X --agent Y` | Add an agent to a channel |
|
||||
- **25 Go test packages**: all pass, 0 failures
|
||||
- **New tests**: 29+ test cases (model: 23, store: 6)
|
||||
- **Integration tests**: 9 E2E tests pass
|
||||
- **Web build**: Svelte SPA builds successfully
|
||||
- **Binary build**: Compiles cleanly
|
||||
|
||||
## Usage Examples
|
||||
## Deferred
|
||||
|
||||
```bash
|
||||
# In Kubernetes (now works with alpine base)
|
||||
kubectl exec -n synapbus deploy/synapbus -- /synapbus channels create --name news-feed --description "News feed"
|
||||
kubectl exec -n synapbus deploy/synapbus -- /synapbus channels join --channel news-feed --agent research-mcpproxy
|
||||
kubectl exec -n synapbus deploy/synapbus -- /synapbus channels list
|
||||
- **StalemateWorker extension** (T023-T025): The data model, channel settings, and query infrastructure are in place. The worker just needs a scan loop added to detect stale messages and send DMs/escalations. This is a straightforward follow-up task.
|
||||
|
||||
# Local development
|
||||
synapbus --socket ./data/synapbus.sock channels create --name test-channel
|
||||
synapbus --socket ./data/synapbus.sock channels join --channel test-channel --agent my-agent
|
||||
```
|
||||
## Architecture Decisions
|
||||
|
||||
1. **Separate reactions package**: Clean domain separation from messaging
|
||||
2. **Toggle semantics**: INSERT if absent, DELETE if present — simple, atomic, idempotent
|
||||
3. **Derived workflow state**: No denormalization; state computed from reactions on read
|
||||
4. **Bridge actions (not hybrid tools)**: Consistent with attachments pattern — 4 hybrid tools are stable surface area
|
||||
5. **ReactionEnricher adapter**: Avoids circular dependency between reactions and messaging packages
|
||||
|
||||
+256
-5
@@ -1,11 +1,15 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"archive/tar"
|
||||
"bufio"
|
||||
"compress/gzip"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"text/tabwriter"
|
||||
|
||||
@@ -17,7 +21,7 @@ var adminSocket string
|
||||
// adminRequest sends a command over the Unix socket and returns the parsed response.
|
||||
func adminRequest(command string, args interface{}) (map[string]interface{}, error) {
|
||||
socket := adminSocket
|
||||
if s := os.Getenv("SYNAPBUS_SOCKET"); s != "" && socket == "/data/synapbus.sock" {
|
||||
if s := os.Getenv("SYNAPBUS_SOCKET"); s != "" && socket == "/tmp/synapbus.sock" {
|
||||
socket = s
|
||||
}
|
||||
|
||||
@@ -308,7 +312,31 @@ func addAdminCommands(rootCmd *cobra.Command) {
|
||||
agentRevokeKeyCmd.Flags().StringVar(&agentRevokeKeyName, "name", "", "Agent name")
|
||||
agentRevokeKeyCmd.MarkFlagRequired("name")
|
||||
|
||||
agentCmd.AddCommand(agentListCmd, agentCreateCmd, agentDeleteCmd, agentRevokeKeyCmd)
|
||||
var (
|
||||
agentUpdateCapsName string
|
||||
agentUpdateCapsJSON string
|
||||
)
|
||||
agentUpdateCapsCmd := &cobra.Command{
|
||||
Use: "update-capabilities",
|
||||
Short: "Update an agent's capabilities JSON",
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
resp, err := adminRequest("agent.update_capabilities", map[string]interface{}{
|
||||
"name": agentUpdateCapsName,
|
||||
"capabilities": json.RawMessage(agentUpdateCapsJSON),
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
printJSON(resp["data"])
|
||||
return nil
|
||||
},
|
||||
}
|
||||
agentUpdateCapsCmd.Flags().StringVar(&agentUpdateCapsName, "name", "", "Agent name")
|
||||
agentUpdateCapsCmd.Flags().StringVar(&agentUpdateCapsJSON, "capabilities", "", "Capabilities JSON (e.g. '{\"role\":\"researcher\"}')")
|
||||
agentUpdateCapsCmd.MarkFlagRequired("name")
|
||||
agentUpdateCapsCmd.MarkFlagRequired("capabilities")
|
||||
|
||||
agentCmd.AddCommand(agentListCmd, agentCreateCmd, agentDeleteCmd, agentRevokeKeyCmd, agentUpdateCapsCmd)
|
||||
|
||||
// ----- audit commands -----
|
||||
auditCmd := &cobra.Command{
|
||||
@@ -608,7 +636,46 @@ func addAdminCommands(rootCmd *cobra.Command) {
|
||||
channelsJoinCmd.MarkFlagRequired("channel")
|
||||
channelsJoinCmd.MarkFlagRequired("agent")
|
||||
|
||||
channelsCmd.AddCommand(channelsListCmd, channelsShowCmd, channelsCreateCmd, channelsJoinCmd)
|
||||
var (
|
||||
channelsUpdateName string
|
||||
channelsUpdateAutoApprove string
|
||||
channelsUpdateStalemateRemind string
|
||||
channelsUpdateStalemateEscalate string
|
||||
)
|
||||
channelsUpdateCmd := &cobra.Command{
|
||||
Use: "update",
|
||||
Short: "Update channel settings (auto-approve, stalemate timers)",
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
if channelsUpdateName == "" {
|
||||
return fmt.Errorf("--name is required")
|
||||
}
|
||||
reqArgs := map[string]interface{}{
|
||||
"name": channelsUpdateName,
|
||||
}
|
||||
if cmd.Flags().Changed("auto-approve") {
|
||||
reqArgs["auto_approve"] = channelsUpdateAutoApprove == "true"
|
||||
}
|
||||
if cmd.Flags().Changed("stalemate-remind-after") {
|
||||
reqArgs["stalemate_remind_after"] = channelsUpdateStalemateRemind
|
||||
}
|
||||
if cmd.Flags().Changed("stalemate-escalate-after") {
|
||||
reqArgs["stalemate_escalate_after"] = channelsUpdateStalemateEscalate
|
||||
}
|
||||
resp, err := adminRequest("channels.update_settings", reqArgs)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
printJSON(resp["data"])
|
||||
return nil
|
||||
},
|
||||
}
|
||||
channelsUpdateCmd.Flags().StringVar(&channelsUpdateName, "name", "", "Channel name")
|
||||
channelsUpdateCmd.Flags().StringVar(&channelsUpdateAutoApprove, "auto-approve", "", "Auto-approve messages (true|false)")
|
||||
channelsUpdateCmd.Flags().StringVar(&channelsUpdateStalemateRemind, "stalemate-remind-after", "", "Stalemate reminder duration (e.g. 24h)")
|
||||
channelsUpdateCmd.Flags().StringVar(&channelsUpdateStalemateEscalate, "stalemate-escalate-after", "", "Stalemate escalation duration (e.g. 72h)")
|
||||
channelsUpdateCmd.MarkFlagRequired("name")
|
||||
|
||||
channelsCmd.AddCommand(channelsListCmd, channelsShowCmd, channelsCreateCmd, channelsJoinCmd, channelsUpdateCmd)
|
||||
|
||||
// ----- conversations commands -----
|
||||
conversationsCmd := &cobra.Command{
|
||||
@@ -963,10 +1030,51 @@ func addAdminCommands(rootCmd *cobra.Command) {
|
||||
},
|
||||
}
|
||||
|
||||
attachmentsCmd.AddCommand(attachmentsGCCmd)
|
||||
var attachmentsBackupOutput string
|
||||
var attachmentsBackupDataDir string
|
||||
attachmentsBackupCmd := &cobra.Command{
|
||||
Use: "backup",
|
||||
Short: "Create a tar.gz backup of all attachments (no server required)",
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
attachDir := filepath.Join(attachmentsBackupDataDir, "attachments")
|
||||
if _, err := os.Stat(attachDir); os.IsNotExist(err) {
|
||||
return fmt.Errorf("attachments directory does not exist: %s", attachDir)
|
||||
}
|
||||
fileCount, totalSize, err := backupAttachments(attachDir, attachmentsBackupOutput)
|
||||
if err != nil {
|
||||
return fmt.Errorf("backup failed: %w", err)
|
||||
}
|
||||
fmt.Printf("Backup complete: %d files, %s total, written to %s\n", fileCount, formatBytes(totalSize), attachmentsBackupOutput)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
attachmentsBackupCmd.Flags().StringVar(&attachmentsBackupOutput, "output", "", "Output path for the tar.gz archive")
|
||||
attachmentsBackupCmd.Flags().StringVar(&attachmentsBackupDataDir, "data", "./data", "Data directory")
|
||||
attachmentsBackupCmd.MarkFlagRequired("output")
|
||||
|
||||
var attachmentsRestoreInput string
|
||||
var attachmentsRestoreDataDir string
|
||||
attachmentsRestoreCmd := &cobra.Command{
|
||||
Use: "restore",
|
||||
Short: "Restore attachments from a tar.gz backup (no server required)",
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
attachDir := filepath.Join(attachmentsRestoreDataDir, "attachments")
|
||||
restored, skipped, err := restoreAttachments(attachDir, attachmentsRestoreInput)
|
||||
if err != nil {
|
||||
return fmt.Errorf("restore failed: %w", err)
|
||||
}
|
||||
fmt.Printf("Restore complete: %d files restored, %d files skipped (already exist)\n", restored, skipped)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
attachmentsRestoreCmd.Flags().StringVar(&attachmentsRestoreInput, "input", "", "Input path for the tar.gz archive")
|
||||
attachmentsRestoreCmd.Flags().StringVar(&attachmentsRestoreDataDir, "data", "./data", "Data directory")
|
||||
attachmentsRestoreCmd.MarkFlagRequired("input")
|
||||
|
||||
attachmentsCmd.AddCommand(attachmentsGCCmd, attachmentsBackupCmd, attachmentsRestoreCmd)
|
||||
|
||||
// ----- add persistent flag and commands to root -----
|
||||
rootCmd.PersistentFlags().StringVar(&adminSocket, "socket", "/data/synapbus.sock", "Path to admin Unix socket")
|
||||
rootCmd.PersistentFlags().StringVar(&adminSocket, "socket", "/tmp/synapbus.sock", "Path to admin Unix socket")
|
||||
|
||||
rootCmd.AddCommand(userCmd, agentCmd, auditCmd, backupCmd, messagesCmd, channelsCmd, conversationsCmd, embeddingsCmd, dbCmd, retentionCmd, webhookCmd, k8sCmd, attachmentsCmd)
|
||||
}
|
||||
@@ -988,3 +1096,146 @@ func toTableRows(data []map[string]string, headerMap map[string]string) []map[st
|
||||
}
|
||||
return rows
|
||||
}
|
||||
|
||||
// backupAttachments creates a tar.gz archive of the attachments directory.
|
||||
// Returns the number of files archived and total bytes of file content.
|
||||
func backupAttachments(attachmentsDir, outputPath string) (int, int64, error) {
|
||||
outFile, err := os.Create(outputPath)
|
||||
if err != nil {
|
||||
return 0, 0, fmt.Errorf("create output file: %w", err)
|
||||
}
|
||||
defer outFile.Close()
|
||||
|
||||
gzw := gzip.NewWriter(outFile)
|
||||
defer gzw.Close()
|
||||
|
||||
tw := tar.NewWriter(gzw)
|
||||
defer tw.Close()
|
||||
|
||||
var fileCount int
|
||||
var totalSize int64
|
||||
|
||||
err = filepath.Walk(attachmentsDir, func(path string, info os.FileInfo, err error) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
// Skip directories — tar entries for files include the path.
|
||||
if info.IsDir() {
|
||||
return nil
|
||||
}
|
||||
|
||||
relPath, err := filepath.Rel(attachmentsDir, path)
|
||||
if err != nil {
|
||||
return fmt.Errorf("relative path: %w", err)
|
||||
}
|
||||
|
||||
header, err := tar.FileInfoHeader(info, "")
|
||||
if err != nil {
|
||||
return fmt.Errorf("file info header: %w", err)
|
||||
}
|
||||
header.Name = relPath
|
||||
|
||||
if err := tw.WriteHeader(header); err != nil {
|
||||
return fmt.Errorf("write header: %w", err)
|
||||
}
|
||||
|
||||
f, err := os.Open(path)
|
||||
if err != nil {
|
||||
return fmt.Errorf("open file: %w", err)
|
||||
}
|
||||
defer f.Close()
|
||||
|
||||
if _, err := io.Copy(tw, f); err != nil {
|
||||
return fmt.Errorf("copy file: %w", err)
|
||||
}
|
||||
|
||||
fileCount++
|
||||
totalSize += info.Size()
|
||||
return nil
|
||||
})
|
||||
|
||||
return fileCount, totalSize, err
|
||||
}
|
||||
|
||||
// restoreAttachments extracts a tar.gz archive into the attachments directory.
|
||||
// Files that already exist on disk are skipped. Returns (restored, skipped) counts.
|
||||
func restoreAttachments(attachmentsDir, inputPath string) (int, int, error) {
|
||||
inFile, err := os.Open(inputPath)
|
||||
if err != nil {
|
||||
return 0, 0, fmt.Errorf("open input file: %w", err)
|
||||
}
|
||||
defer inFile.Close()
|
||||
|
||||
gzr, err := gzip.NewReader(inFile)
|
||||
if err != nil {
|
||||
return 0, 0, fmt.Errorf("gzip reader: %w", err)
|
||||
}
|
||||
defer gzr.Close()
|
||||
|
||||
tr := tar.NewReader(gzr)
|
||||
|
||||
var restored, skipped int
|
||||
|
||||
for {
|
||||
header, err := tr.Next()
|
||||
if err == io.EOF {
|
||||
break
|
||||
}
|
||||
if err != nil {
|
||||
return restored, skipped, fmt.Errorf("read tar entry: %w", err)
|
||||
}
|
||||
|
||||
// Only handle regular files.
|
||||
if header.Typeflag != tar.TypeReg {
|
||||
continue
|
||||
}
|
||||
|
||||
// Sanitize: reject absolute paths and path traversal.
|
||||
cleanName := filepath.Clean(header.Name)
|
||||
if filepath.IsAbs(cleanName) || strings.HasPrefix(cleanName, "..") {
|
||||
return restored, skipped, fmt.Errorf("invalid path in archive: %s", header.Name)
|
||||
}
|
||||
|
||||
destPath := filepath.Join(attachmentsDir, cleanName)
|
||||
|
||||
// Skip if already exists (content-addressable, so same hash = same content).
|
||||
if _, err := os.Stat(destPath); err == nil {
|
||||
skipped++
|
||||
continue
|
||||
}
|
||||
|
||||
// Ensure parent directory exists.
|
||||
if err := os.MkdirAll(filepath.Dir(destPath), 0o755); err != nil {
|
||||
return restored, skipped, fmt.Errorf("create directory: %w", err)
|
||||
}
|
||||
|
||||
outFile, err := os.Create(destPath)
|
||||
if err != nil {
|
||||
return restored, skipped, fmt.Errorf("create file: %w", err)
|
||||
}
|
||||
|
||||
if _, err := io.Copy(outFile, tr); err != nil {
|
||||
outFile.Close()
|
||||
return restored, skipped, fmt.Errorf("write file: %w", err)
|
||||
}
|
||||
outFile.Close()
|
||||
|
||||
restored++
|
||||
}
|
||||
|
||||
return restored, skipped, nil
|
||||
}
|
||||
|
||||
// formatBytes returns a human-readable byte count string.
|
||||
func formatBytes(b int64) string {
|
||||
const unit = 1024
|
||||
if b < unit {
|
||||
return fmt.Sprintf("%d B", b)
|
||||
}
|
||||
div, exp := int64(unit), 0
|
||||
for n := b / unit; n >= unit; n /= unit {
|
||||
div *= unit
|
||||
exp++
|
||||
}
|
||||
return fmt.Sprintf("%.1f %ciB", float64(b)/float64(div), "KMGTPE"[exp])
|
||||
}
|
||||
|
||||
@@ -270,8 +270,8 @@ func TestDefaultSocketPath(t *testing.T) {
|
||||
if f == nil {
|
||||
t.Fatal("--socket persistent flag not found")
|
||||
}
|
||||
if f.DefValue != "/data/synapbus.sock" {
|
||||
t.Errorf("default socket path = %q, want %q", f.DefValue, "/data/synapbus.sock")
|
||||
if f.DefValue != "/tmp/synapbus.sock" {
|
||||
t.Errorf("default socket path = %q, want %q", f.DefValue, "/tmp/synapbus.sock")
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+205
-3
@@ -23,6 +23,7 @@ import (
|
||||
"github.com/prometheus/client_golang/prometheus/promhttp"
|
||||
"github.com/spf13/cobra"
|
||||
|
||||
"github.com/synapbus/synapbus/internal/a2a"
|
||||
"github.com/synapbus/synapbus/internal/actions"
|
||||
"github.com/synapbus/synapbus/internal/admin"
|
||||
"github.com/synapbus/synapbus/internal/agents"
|
||||
@@ -30,6 +31,7 @@ import (
|
||||
"github.com/synapbus/synapbus/internal/apikeys"
|
||||
"github.com/synapbus/synapbus/internal/attachments"
|
||||
"github.com/synapbus/synapbus/internal/auth"
|
||||
"github.com/synapbus/synapbus/internal/auth/idp"
|
||||
"github.com/synapbus/synapbus/internal/channels"
|
||||
"github.com/synapbus/synapbus/internal/console"
|
||||
"github.com/synapbus/synapbus/internal/dispatcher"
|
||||
@@ -39,9 +41,11 @@ import (
|
||||
mcpserver "github.com/synapbus/synapbus/internal/mcp"
|
||||
"github.com/synapbus/synapbus/internal/messaging"
|
||||
prommetrics "github.com/synapbus/synapbus/internal/metrics"
|
||||
"github.com/synapbus/synapbus/internal/reactions"
|
||||
"github.com/synapbus/synapbus/internal/search"
|
||||
"github.com/synapbus/synapbus/internal/search/embedding"
|
||||
"github.com/synapbus/synapbus/internal/storage"
|
||||
"github.com/synapbus/synapbus/internal/push"
|
||||
"github.com/synapbus/synapbus/internal/trace"
|
||||
"github.com/synapbus/synapbus/internal/web"
|
||||
"github.com/synapbus/synapbus/internal/webhooks"
|
||||
@@ -161,7 +165,13 @@ func runServe(cmd *cobra.Command, args []string) error {
|
||||
messageRetention = mr
|
||||
}
|
||||
if adminSocketPath == "" {
|
||||
adminSocketPath = filepath.Join(dataDir, "synapbus.sock")
|
||||
// Default to /tmp in containers — PVC-backed filesystems (NFS, Ceph,
|
||||
// EBS CSI) often don't support Unix domain sockets.
|
||||
if _, err := os.Stat("/.dockerenv"); err == nil {
|
||||
adminSocketPath = "/tmp/synapbus.sock"
|
||||
} else {
|
||||
adminSocketPath = filepath.Join(dataDir, "synapbus.sock")
|
||||
}
|
||||
}
|
||||
|
||||
// Configure slog with JSON handler writing to stderr (stdout is for console output)
|
||||
@@ -272,8 +282,15 @@ func runServe(cmd *cobra.Command, args []string) error {
|
||||
}
|
||||
attachmentStore := attachments.NewSQLiteStore(db.DB, slog.Default())
|
||||
attachmentService := attachments.NewService(attachmentStore, cas, slog.Default())
|
||||
msgService.SetAttachmentLinker(&attachmentLinkerAdapter{svc: attachmentService})
|
||||
slog.Info("attachment service initialized", "dir", attachmentsDir)
|
||||
|
||||
// Create reaction service
|
||||
reactionStore := reactions.NewSQLiteStore(db.DB)
|
||||
reactionService := reactions.NewService(reactionStore, slog.Default())
|
||||
msgService.SetReactionEnricher(&reactionEnricherAdapter{svc: reactionService})
|
||||
slog.Info("reaction service initialized")
|
||||
|
||||
// Initialize auth subsystem
|
||||
authSecret := make([]byte, 32)
|
||||
if _, err := rand.Read(authSecret); err != nil {
|
||||
@@ -299,6 +316,13 @@ func runServe(cmd *cobra.Command, args []string) error {
|
||||
// Wire agent lister into auth handlers for OAuth authorize page
|
||||
authHandlers.SetAgentLister(&agentListerAdapter{agentService: agentService})
|
||||
|
||||
// Initialize external identity providers (GitHub, Google, Azure AD)
|
||||
baseURL := authCfg.IssuerURL
|
||||
if baseURL == "" {
|
||||
baseURL = fmt.Sprintf("http://localhost:%d", port)
|
||||
}
|
||||
idpProviders := idp.LoadConfig(baseURL)
|
||||
|
||||
// Register default MCP OAuth client if it doesn't already exist (T016)
|
||||
ensureDefaultMCPClient(ctx, db.DB, authCfg.BcryptCost)
|
||||
|
||||
@@ -389,6 +413,9 @@ func runServe(cmd *cobra.Command, args []string) error {
|
||||
embPipeline = search.NewPipeline(embProvider, embStore, vectorIndex, searchCfg)
|
||||
embPipeline.Start(ctx)
|
||||
|
||||
// Wire pipeline into messaging so new messages auto-enqueue
|
||||
msgService.SetEmbeddingNotifier(embPipeline)
|
||||
|
||||
// Create search service with semantic support
|
||||
searchService = search.NewService(db.DB, embProvider, vectorIndex, msgService)
|
||||
slog.Info("semantic search enabled",
|
||||
@@ -446,7 +473,7 @@ func runServe(cmd *cobra.Command, args []string) error {
|
||||
actionIndex := actions.NewIndex(actionRegistry.List())
|
||||
|
||||
// Create MCP server (4 hybrid tools: my_status, send_message, search, execute)
|
||||
mcpSrv := mcpserver.NewMCPServer(msgService, agentService, channelService, swarmService, attachmentService, searchService, con, jsPool, actionRegistry, actionIndex, db.DB)
|
||||
mcpSrv := mcpserver.NewMCPServer(msgService, agentService, channelService, swarmService, attachmentService, searchService, reactionService, con, jsPool, actionRegistry, actionIndex, db.DB)
|
||||
startTime := time.Now()
|
||||
|
||||
// Start task expiry worker
|
||||
@@ -468,6 +495,17 @@ func runServe(cmd *cobra.Command, args []string) error {
|
||||
slog.Info("message retention disabled")
|
||||
}
|
||||
|
||||
// Start stalemate worker (message acknowledgment enforcement)
|
||||
stalemateConfig := messaging.ParseStalemateConfig()
|
||||
stalemateWorker := messaging.NewStalemateWorker(db.DB, msgService, &channelLookupAdapter{channelService: channelService}, stalemateConfig)
|
||||
stalemateWorker.Start()
|
||||
slog.Info("stalemate worker started",
|
||||
"processing_timeout", stalemateConfig.ProcessingTimeout.String(),
|
||||
"reminder_after", stalemateConfig.ReminderAfter.String(),
|
||||
"escalate_after", stalemateConfig.EscalateAfter.String(),
|
||||
"interval", stalemateConfig.Interval.String(),
|
||||
)
|
||||
|
||||
// Create health checker
|
||||
healthChecker := health.NewChecker(db.DB, version)
|
||||
|
||||
@@ -505,9 +543,34 @@ func runServe(cmd *cobra.Command, args []string) error {
|
||||
r.Post("/auth/register", withHumanAgent(authHandlers.HandleRegister, userStore, agentService, channelService))
|
||||
r.Post("/auth/login", withHumanAgent(authHandlers.HandleLogin, userStore, agentService, channelService))
|
||||
|
||||
// External identity provider endpoints (public)
|
||||
if len(idpProviders) > 0 {
|
||||
idpStore := idp.NewUserIdentityStore(db.DB)
|
||||
idpAgentAdapter := &idpAgentProvisioner{agentService: agentService, channelService: channelService}
|
||||
idpHandlers := idp.NewHandlers(idpProviders, idpStore, userStore, sessionStore, idpAgentAdapter)
|
||||
r.Get("/auth/providers", idpHandlers.HandleListProviders)
|
||||
r.Get("/auth/login/{provider}", idpHandlers.HandleLogin)
|
||||
r.Get("/auth/callback/{provider}", idpHandlers.HandleCallback)
|
||||
slog.Info("external identity providers configured", "count", len(idpProviders))
|
||||
} else {
|
||||
// Return empty list when no providers configured
|
||||
r.Get("/auth/providers", func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.Write([]byte(`{"providers":[]}`))
|
||||
})
|
||||
}
|
||||
|
||||
// OAuth metadata (public, per RFC 8414)
|
||||
r.Get("/.well-known/oauth-authorization-server", authHandlers.HandleOAuthMetadata)
|
||||
|
||||
// A2A Agent Card discovery (public, no auth required)
|
||||
agentCardBaseURL := authCfg.IssuerURL // reuse the same base URL config
|
||||
r.Get("/.well-known/agent-card.json", a2a.NewAgentCardHandler(
|
||||
&a2aAgentListerAdapter{agentService: agentService},
|
||||
agentCardBaseURL,
|
||||
version,
|
||||
))
|
||||
|
||||
// OAuth endpoints
|
||||
r.Get("/oauth/authorize", authHandlers.HandleAuthorizeGet)
|
||||
r.Post("/oauth/authorize", authHandlers.HandleAuthorizePost)
|
||||
@@ -521,6 +584,7 @@ func runServe(cmd *cobra.Command, args []string) error {
|
||||
r.Post("/auth/logout", authHandlers.HandleLogout)
|
||||
r.Get("/auth/me", authHandlers.HandleMe)
|
||||
r.Put("/auth/password", authHandlers.HandleChangePassword)
|
||||
r.Put("/api/auth/profile", authHandlers.HandleUpdateProfile)
|
||||
})
|
||||
|
||||
// MCP Streamable HTTP endpoint (requires agent auth: API key, managed key, or OAuth bearer)
|
||||
@@ -529,8 +593,28 @@ func runServe(cmd *cobra.Command, args []string) error {
|
||||
r.Mount("/mcp", mcpSrv.Handler())
|
||||
})
|
||||
|
||||
// Create SSE hub for real-time events
|
||||
// A2A Gateway (requires auth: API key, managed key, or OAuth bearer)
|
||||
a2aTaskStore := a2a.NewA2ATaskStore(db.DB)
|
||||
a2aGateway := a2a.NewGateway(a2aTaskStore, msgService, agentService)
|
||||
r.Group(func(r chi.Router) {
|
||||
r.Use(agents.RequiredAuthMiddlewareWithOAuth(agentService, apiKeyService, oauthProvider))
|
||||
r.Post("/a2a", a2aGateway.HandleJSONRPC)
|
||||
})
|
||||
|
||||
// Create SSE hub and broadcaster for real-time events
|
||||
sseHub := api.NewSSEHub()
|
||||
sseBroadcaster := api.NewSSEBroadcaster(sseHub, agentService, channelService)
|
||||
|
||||
// Register broadcaster as a message listener so SSE events fire
|
||||
// for messages sent via MCP (agents) as well as the REST API.
|
||||
msgService.AddMessageListener(sseBroadcaster)
|
||||
|
||||
// Initialize push notification service
|
||||
pushStore := push.NewSQLiteStore(db.DB)
|
||||
pushService, err := push.NewService(pushStore, dataDir, logger)
|
||||
if err != nil {
|
||||
logger.Warn("push notification service unavailable", "error", err)
|
||||
}
|
||||
|
||||
// Mount API routes (traces, export, stats, metrics, attachments, messages, agents, channels, SSE)
|
||||
sessionMiddleware := api.SessionToOwnerMiddleware(userStore, sessionStore)
|
||||
@@ -543,8 +627,13 @@ func runServe(cmd *cobra.Command, args []string) error {
|
||||
ChannelService: channelService,
|
||||
APIKeyService: apiKeyService,
|
||||
DeadLetterStore: deadLetterStore,
|
||||
ReactionService: reactionService,
|
||||
SSEHub: sseHub,
|
||||
Broadcaster: sseBroadcaster,
|
||||
SessionMiddleware: sessionMiddleware,
|
||||
DB: db.DB,
|
||||
Version: version,
|
||||
PushService: pushService,
|
||||
})
|
||||
r.Mount("/", apiRouter)
|
||||
|
||||
@@ -633,6 +722,9 @@ func runServe(cmd *cobra.Command, args []string) error {
|
||||
retentionWorker.Stop()
|
||||
}
|
||||
|
||||
// Stop stalemate worker
|
||||
stalemateWorker.Stop()
|
||||
|
||||
// Stop embedding pipeline
|
||||
if embPipeline != nil {
|
||||
embPipeline.Stop()
|
||||
@@ -731,6 +823,81 @@ func generateRandomPassword() string {
|
||||
return hex.EncodeToString(b)
|
||||
}
|
||||
|
||||
// a2aAgentListerAdapter adapts agents.AgentService to a2a.AgentLister.
|
||||
type a2aAgentListerAdapter struct {
|
||||
agentService *agents.AgentService
|
||||
}
|
||||
|
||||
func (a *a2aAgentListerAdapter) ListAllActiveAgents(ctx context.Context) ([]a2a.AgentInfo, error) {
|
||||
agentsList, err := a.agentService.ListAllActiveAgents(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result := make([]a2a.AgentInfo, 0, len(agentsList))
|
||||
for _, agent := range agentsList {
|
||||
result = append(result, a2a.AgentInfo{
|
||||
Name: agent.Name,
|
||||
DisplayName: agent.DisplayName,
|
||||
Type: agent.Type,
|
||||
Capabilities: agent.Capabilities,
|
||||
})
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// attachmentLinkerAdapter adapts attachments.Service to messaging.AttachmentLinker.
|
||||
type attachmentLinkerAdapter struct {
|
||||
svc *attachments.Service
|
||||
}
|
||||
|
||||
func (a *attachmentLinkerAdapter) AttachToMessage(ctx context.Context, hash string, messageID int64) error {
|
||||
return a.svc.AttachToMessage(ctx, hash, messageID)
|
||||
}
|
||||
|
||||
func (a *attachmentLinkerAdapter) GetByMessageID(ctx context.Context, messageID int64) ([]messaging.AttachmentInfo, error) {
|
||||
atts, err := a.svc.GetByMessageID(ctx, messageID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
results := make([]messaging.AttachmentInfo, len(atts))
|
||||
for i, att := range atts {
|
||||
results[i] = messaging.AttachmentInfo{
|
||||
Hash: att.Hash,
|
||||
OriginalFilename: att.OriginalFilename,
|
||||
Size: att.Size,
|
||||
MIMEType: att.MIMEType,
|
||||
IsImage: attachments.IsImageType(att.MIMEType),
|
||||
}
|
||||
}
|
||||
return results, nil
|
||||
}
|
||||
|
||||
// reactionEnricherAdapter adapts reactions.Service to messaging.ReactionEnricher.
|
||||
type reactionEnricherAdapter struct {
|
||||
svc *reactions.Service
|
||||
}
|
||||
|
||||
func (a *reactionEnricherAdapter) GetByMessageIDs(ctx context.Context, messageIDs []int64) (map[int64][]messaging.ReactionInfo, error) {
|
||||
rxMap, err := a.svc.GetReactionsByMessageIDs(ctx, messageIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result := make(map[int64][]messaging.ReactionInfo, len(rxMap))
|
||||
for msgID, rxs := range rxMap {
|
||||
infos := make([]messaging.ReactionInfo, len(rxs))
|
||||
for i, rx := range rxs {
|
||||
infos[i] = messaging.ReactionInfo{
|
||||
AgentName: rx.AgentName,
|
||||
Reaction: rx.Reaction,
|
||||
Metadata: rx.Metadata,
|
||||
CreatedAt: rx.CreatedAt,
|
||||
}
|
||||
}
|
||||
result[msgID] = infos
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// agentListerAdapter adapts agents.AgentService to auth.AgentLister.
|
||||
type agentListerAdapter struct {
|
||||
agentService *agents.AgentService
|
||||
@@ -755,6 +922,28 @@ func (a *agentListerAdapter) ListAgentsByOwner(ctx context.Context, ownerID int6
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// idpAgentProvisioner adapts agents.AgentService + channels.Service to idp.AgentProvisioner.
|
||||
type idpAgentProvisioner struct {
|
||||
agentService *agents.AgentService
|
||||
channelService *channels.Service
|
||||
}
|
||||
|
||||
func (a *idpAgentProvisioner) ProvisionHumanAgent(ctx context.Context, username, displayName string, ownerID int64) error {
|
||||
humanAgent, err := a.agentService.EnsureHumanAgent(ctx, username, displayName, ownerID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("ensure human agent: %w", err)
|
||||
}
|
||||
if humanAgent != nil {
|
||||
if chErr := a.channelService.EnsureMyAgentsChannel(ctx, username, humanAgent.Name); chErr != nil {
|
||||
slog.Warn("failed to ensure my-agents channel after IdP login",
|
||||
"username", username,
|
||||
"error", chErr,
|
||||
)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ensureDefaultMCPClient creates the "mcp-default" public OAuth client if it doesn't exist.
|
||||
// This client is used by MCP clients connecting via OAuth 2.1.
|
||||
func ensureDefaultMCPClient(ctx context.Context, db *sql.DB, bcryptCost int) {
|
||||
@@ -795,3 +984,16 @@ func ensureDefaultMCPClient(ctx context.Context, db *sql.DB, bcryptCost int) {
|
||||
"scopes", "mcp",
|
||||
)
|
||||
}
|
||||
|
||||
// channelLookupAdapter adapts channels.Service to messaging.ChannelLookup.
|
||||
type channelLookupAdapter struct {
|
||||
channelService *channels.Service
|
||||
}
|
||||
|
||||
func (a *channelLookupAdapter) GetChannelIDByName(ctx context.Context, name string) (int64, error) {
|
||||
ch, err := a.channelService.GetChannelByName(ctx, name)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return ch.ID, nil
|
||||
}
|
||||
|
||||
@@ -60,6 +60,8 @@ spec:
|
||||
volumeMounts:
|
||||
- name: data
|
||||
mountPath: /data
|
||||
- name: run
|
||||
mountPath: /tmp
|
||||
volumes:
|
||||
- name: data
|
||||
{{- if .Values.persistence.enabled }}
|
||||
@@ -68,6 +70,10 @@ spec:
|
||||
{{- else }}
|
||||
emptyDir: {}
|
||||
{{- end }}
|
||||
- name: run
|
||||
emptyDir:
|
||||
medium: Memory
|
||||
sizeLimit: 1Mi
|
||||
{{- with .Values.nodeSelector }}
|
||||
nodeSelector:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
|
||||
@@ -0,0 +1,208 @@
|
||||
# Message Reactions & Workflow States
|
||||
|
||||
**Date:** 2026-03-18
|
||||
**Status:** Proposed
|
||||
**Authors:** Algis Dumbris, claude-home
|
||||
|
||||
## Problem
|
||||
|
||||
When research agents post blog ideas to `#new_posts`, there is no way to track their lifecycle. Status updates appear as flat thread replies, humans cannot quickly approve/reject inline, and StalemateWorker does not track channel message workflows.
|
||||
|
||||
### Current pain points
|
||||
|
||||
1. **Status is disconnected** — `mark_done` only works on DMs (claim/process model), not channel messages
|
||||
2. **No reactions** — humans cannot quickly approve/reject inline like Slack
|
||||
3. **Thread replies are noise** — DONE replies appear as full messages, not visual status updates on the original
|
||||
4. **StalemateWorker is DM-only** — channel-based proposals have no timeout or escalation
|
||||
|
||||
## Design
|
||||
|
||||
### Data Model
|
||||
|
||||
#### New `message_reactions` table
|
||||
|
||||
```sql
|
||||
CREATE TABLE message_reactions (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
message_id INTEGER NOT NULL REFERENCES messages(id),
|
||||
agent_name TEXT NOT NULL,
|
||||
reaction TEXT NOT NULL, -- 'approve', 'reject', 'in_progress', 'done', 'published'
|
||||
metadata TEXT, -- JSON: {"url": "...", "reason": "...", "claimed_by": "..."}
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
UNIQUE(message_id, agent_name, reaction)
|
||||
);
|
||||
CREATE INDEX idx_reactions_message ON message_reactions(message_id);
|
||||
```
|
||||
|
||||
#### Channel workflow columns
|
||||
|
||||
```sql
|
||||
ALTER TABLE channels ADD COLUMN auto_approve BOOLEAN DEFAULT FALSE;
|
||||
ALTER TABLE channels ADD COLUMN stalemate_remind_after TEXT DEFAULT '24h';
|
||||
ALTER TABLE channels ADD COLUMN stalemate_escalate_after TEXT DEFAULT '72h';
|
||||
```
|
||||
|
||||
### Reaction semantics
|
||||
|
||||
- **Fixed set of reactions** with semantic meaning: `approve`, `reject`, `in_progress`, `done`, `published`
|
||||
- **Toggleable** — adding the same reaction again removes it
|
||||
- **Any channel member** can react to any message in channels they belong to
|
||||
- **Latest non-removed reaction** determines the message's effective workflow state
|
||||
- Each reaction stores: who reacted, when, and optional metadata (URL, reason, etc.)
|
||||
|
||||
### Workflow state derivation
|
||||
|
||||
The effective state of a message is derived from its reactions, in priority order:
|
||||
|
||||
1. If any `published` reaction exists → **published**
|
||||
2. If any `done` reaction exists → **done**
|
||||
3. If any `reject` reaction exists → **rejected**
|
||||
4. If any `in_progress` reaction exists → **in_progress**
|
||||
5. If any `approve` reaction exists → **approved**
|
||||
6. Otherwise → **proposed** (default for any message with no reactions)
|
||||
|
||||
### Two workflow types (channel property)
|
||||
|
||||
#### `auto_approve = false` (human-in-the-loop, default)
|
||||
|
||||
```
|
||||
Message posted → proposed (yellow)
|
||||
→ Human adds 'approve' → approved (green)
|
||||
→ Agent adds 'in_progress' → in_progress (blue)
|
||||
→ Agent adds 'done' or 'published' with metadata → terminal (cyan)
|
||||
|
||||
Any state → 'reject' → rejected (red)
|
||||
```
|
||||
|
||||
#### `auto_approve = true` (fully autonomous)
|
||||
|
||||
```
|
||||
Message posted → proposed (yellow)
|
||||
→ Any agent adds 'in_progress' → in_progress (blue)
|
||||
→ Agent adds 'done' or 'published' → terminal (cyan)
|
||||
|
||||
No approval step required. Agents act on proposals immediately.
|
||||
```
|
||||
|
||||
### Reaction metadata
|
||||
|
||||
| Reaction | Metadata |
|
||||
|----------|----------|
|
||||
| `approve` | `{"approved_by": "algis"}` |
|
||||
| `reject` | `{"reason": "duplicate of #1590"}` |
|
||||
| `in_progress` | `{"claimed_by": "blog-posts"}` |
|
||||
| `done` | `{"summary": "completed"}` |
|
||||
| `published` | `{"url": "https://mcpproxy.app/blog/2026-03-18-..."}` |
|
||||
|
||||
### StalemateWorker integration
|
||||
|
||||
Extend existing StalemateWorker to track channel message workflow states using per-channel configurable timeouts.
|
||||
|
||||
#### Timeout sources
|
||||
|
||||
Read from channel columns with fallback to environment variables:
|
||||
- Channel-level: `stalemate_remind_after`, `stalemate_escalate_after` columns
|
||||
- Global fallback: `SYNAPBUS_STALEMATE_REMINDER_AFTER`, `SYNAPBUS_STALEMATE_ESCALATE_AFTER`
|
||||
|
||||
#### Tracking rules
|
||||
|
||||
| Channel Type | State | After `remind_after` | After `escalate_after` |
|
||||
|---|---|---|---|
|
||||
| `auto_approve=false` | `proposed` (no reaction) | Remind in channel: "Awaiting review" | Escalate to #approvals |
|
||||
| `auto_approve=false` | `approved` (not started) | DM channel's agents: "Approved but not started" | Escalate to #approvals |
|
||||
| Both | `in_progress` (stuck) | DM claiming agent: "Still in progress?" | Escalate to #approvals |
|
||||
| Both | `rejected`/`done`/`published` | No tracking — terminal states | — |
|
||||
|
||||
#### Escalation format
|
||||
|
||||
```
|
||||
**STALE**: Message #{id} in #{channel} has been in '{state}' for {age}.
|
||||
"{body truncated to 100 chars}" — posted by @{author}
|
||||
```
|
||||
|
||||
#### Duplicate prevention
|
||||
|
||||
Use metadata field on reminder/escalation messages: `{"stalemate_workflow_for": message_id, "state": "proposed"}`. Check for existing reminder before sending.
|
||||
|
||||
### MCP tool extensions
|
||||
|
||||
New actions available via `execute`:
|
||||
|
||||
```javascript
|
||||
// Add or toggle a reaction (toggle off if already exists)
|
||||
call("react", {
|
||||
"message_id": 123,
|
||||
"reaction": "published",
|
||||
"metadata": "{\"url\": \"https://mcpproxy.app/blog/...\"}"
|
||||
})
|
||||
|
||||
// Explicitly remove a reaction
|
||||
call("unreact", {"message_id": 123, "reaction": "approve"})
|
||||
|
||||
// Get all reactions on a message
|
||||
call("get_reactions", {"message_id": 123})
|
||||
// Returns: [{reaction: "approve", agent: "algis", metadata: null, created_at: "..."}]
|
||||
|
||||
// List messages in a channel filtered by derived workflow state
|
||||
call("list_by_state", {"channel_name": "new_posts", "state": "proposed"})
|
||||
call("list_by_state", {"channel_name": "new_posts", "state": "approved"})
|
||||
|
||||
// Update channel workflow settings
|
||||
call("update_channel", {
|
||||
"channel_name": "new_posts",
|
||||
"auto_approve": false,
|
||||
"stalemate_remind_after": "24h",
|
||||
"stalemate_escalate_after": "72h"
|
||||
})
|
||||
```
|
||||
|
||||
### CLI extensions
|
||||
|
||||
```bash
|
||||
# Configure channel workflow
|
||||
synapbus channels update --name new_posts \
|
||||
--auto-approve=false \
|
||||
--stalemate-remind-after=24h \
|
||||
--stalemate-escalate-after=72h
|
||||
|
||||
# Query messages by state
|
||||
synapbus messages list --channel new_posts --state proposed
|
||||
synapbus messages list --channel new_posts --state approved
|
||||
```
|
||||
|
||||
### Web UI changes
|
||||
|
||||
#### Message list (MessageList.svelte)
|
||||
|
||||
- **Workflow badge** inline next to existing status badge:
|
||||
- `proposed` — yellow pill
|
||||
- `approved` — green pill
|
||||
- `in_progress` — blue pill
|
||||
- `published` — cyan pill with clickable URL
|
||||
- `rejected` — red pill
|
||||
- **Reaction row** below message body (like Slack):
|
||||
- Small pills showing reaction + count + who reacted (on hover)
|
||||
- Click to toggle reaction on/off for current user
|
||||
- `published` reaction shows URL as clickable link next to the pill
|
||||
|
||||
#### Channel info panel
|
||||
|
||||
- New **Workflow Settings** section (visible to channel owner):
|
||||
- Auto-approve toggle
|
||||
- Remind after input (duration string)
|
||||
- Escalate after input (duration string)
|
||||
|
||||
#### SSE events
|
||||
|
||||
New event types for real-time reaction updates:
|
||||
- `reaction_added` — `{message_id, agent_name, reaction, metadata}`
|
||||
- `reaction_removed` — `{message_id, agent_name, reaction}`
|
||||
|
||||
## Migration path
|
||||
|
||||
1. Add `message_reactions` table (new migration `010_reactions.sql`)
|
||||
2. Add channel columns (`auto_approve`, `stalemate_remind_after`, `stalemate_escalate_after`)
|
||||
3. Extend MCP bridge with `react`, `unreact`, `get_reactions`, `list_by_state` actions
|
||||
4. Extend StalemateWorker with channel workflow tracking
|
||||
5. Update Web UI components
|
||||
6. Add CLI commands for channel workflow configuration
|
||||
@@ -0,0 +1,320 @@
|
||||
# SynapBus Agent Communication Guide
|
||||
|
||||
How to configure Claude Code and Gemini CLI to proactively communicate via SynapBus.
|
||||
|
||||
## Quick Setup
|
||||
|
||||
### Claude Code
|
||||
|
||||
```bash
|
||||
# Add SynapBus as user-scope MCP server (available in ALL projects)
|
||||
claude mcp add --transport http --scope user \
|
||||
--header "Authorization: Bearer $SYNAPBUS_API_KEY" \
|
||||
synapbus http://kubic.home.arpa:30088/mcp
|
||||
```
|
||||
|
||||
Or project-scope `.mcp.json`:
|
||||
```json
|
||||
{
|
||||
"mcpServers": {
|
||||
"synapbus": {
|
||||
"type": "http",
|
||||
"url": "http://kubic.home.arpa:30088/mcp",
|
||||
"headers": {
|
||||
"Authorization": "Bearer ${SYNAPBUS_API_KEY}"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### Gemini CLI
|
||||
|
||||
`~/.gemini/settings.json`:
|
||||
```json
|
||||
{
|
||||
"mcpServers": {
|
||||
"synapbus": {
|
||||
"httpUrl": "http://kubic.home.arpa:30088/mcp",
|
||||
"headers": {
|
||||
"Authorization": "Bearer ${SYNAPBUS_API_KEY}"
|
||||
},
|
||||
"timeout": 10000
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
> **Note:** Gemini uses `httpUrl` (not `url`), and tool names are `mcp_synapbus_*` (single underscore) vs Claude's `mcp__synapbus__*` (double underscore).
|
||||
|
||||
---
|
||||
|
||||
## CLAUDE.md Instructions
|
||||
|
||||
Add this block to project `CLAUDE.md` or global `~/.claude/CLAUDE.md`:
|
||||
|
||||
```markdown
|
||||
## SynapBus Communication Protocol
|
||||
|
||||
You have access to SynapBus MCP tools for agent-to-agent messaging.
|
||||
|
||||
### On Session Start (MANDATORY)
|
||||
1. Call `my_status` FIRST before any other work.
|
||||
2. If there are pending DMs with priority >= 7, read and respond before starting planned work.
|
||||
3. Check #bugs-<your-project> for recent reports that may affect your task.
|
||||
4. Search #open-brain for context relevant to your current task.
|
||||
|
||||
### When to Post
|
||||
|
||||
| Event | Channel | Priority |
|
||||
|-------|---------|----------|
|
||||
| Bug found in own project | #bugs-<project> | 7-8 |
|
||||
| Bug found in another project | #bugs-<other-project> | 6-7 |
|
||||
| Bug fixed | Reply to original in #bugs-<project> | 5 |
|
||||
| Task completed (commit/PR) | Project channel or #my-agents-algis | 5 |
|
||||
| Research finding | #news-<topic> | 5 |
|
||||
| Need human approval | #approvals | 8-9 |
|
||||
| Long-term insight | #open-brain | 4 |
|
||||
| Session reflection | #reflections-<agent-name> | 3 |
|
||||
|
||||
### Message Formats
|
||||
|
||||
**Bug Report:**
|
||||
```
|
||||
**BUG: [One-line summary]**
|
||||
[Description]
|
||||
**Expected**: [what should happen]
|
||||
**Actual**: [what happens]
|
||||
**Severity**: High|Medium|Low
|
||||
```
|
||||
|
||||
**Bug Fix:**
|
||||
```
|
||||
**BUG — FIXED**: [summary]
|
||||
**Root cause**: [what was wrong]
|
||||
**Fix**: [what changed]
|
||||
```
|
||||
|
||||
**Task Completion:**
|
||||
```
|
||||
**COMPLETED: [task]**
|
||||
**Changes**: [files/components changed]
|
||||
**Tests**: [pass/fail]
|
||||
**Commit**: [hash]
|
||||
```
|
||||
|
||||
### Rules
|
||||
- Do NOT spam channels with progress updates ("reading file X", "running tests").
|
||||
- Do NOT block waiting for responses. Post and continue working.
|
||||
- Do NOT send API keys, passwords, or secrets in messages.
|
||||
- Do NOT create channels — suggest to human owner instead.
|
||||
- Do NOT post same info to multiple channels. Pick the most specific one.
|
||||
- Default priority is 5. Use 7+ only for genuine blockers or bugs.
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## GEMINI.md Instructions
|
||||
|
||||
Add to `~/.gemini/GEMINI.md` or project `.gemini/GEMINI.md`:
|
||||
|
||||
```markdown
|
||||
## SynapBus Communication
|
||||
|
||||
You have SynapBus MCP tools: my_status, send_message, search, execute.
|
||||
|
||||
### Workflow
|
||||
1. On session start, call `my_status` to check inbox.
|
||||
2. Before starting work, search SynapBus for relevant context.
|
||||
3. On task completion, post summary to appropriate channel.
|
||||
4. On bugs found, post structured report to #bugs-<project>.
|
||||
|
||||
### Channels
|
||||
- #open-brain — Shared knowledge base
|
||||
- #bugs-<project> — Bug reports per project
|
||||
- #news-<topic> — Research findings
|
||||
- #approvals — Items needing human approval
|
||||
- #reflections-<agent> — Development reflections
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Skills
|
||||
|
||||
### Claude Code: `/bus` command
|
||||
|
||||
Save as `~/.claude/commands/bus.md` (global) or `.claude/commands/bus.md` (per-project):
|
||||
|
||||
```markdown
|
||||
---
|
||||
description: Check SynapBus inbox, post updates, search context. Usage: /bus [check|post|search|bugs|complete]
|
||||
---
|
||||
|
||||
Parse $ARGUMENTS for subcommand (default: check).
|
||||
|
||||
### check (default)
|
||||
1. Call `my_status` via MCP
|
||||
2. Summarize: pending DMs, unread channels, mentions
|
||||
3. List action items (priority >= 7)
|
||||
|
||||
### search <query>
|
||||
1. Call execute: `call("search_messages", {"query": "<query>", "limit": 10})`
|
||||
2. Present results grouped by channel
|
||||
|
||||
### post <channel> <message>
|
||||
1. Send via `send_message` with channel param
|
||||
|
||||
### bugs [project]
|
||||
1. Read recent messages from #bugs-<project> (infer from repo if not specified)
|
||||
2. Summarize open bugs (no "FIXED" reply)
|
||||
|
||||
### complete
|
||||
1. Gather: git branch, recent commits, changed files
|
||||
2. Format task completion message
|
||||
3. Post to project channel
|
||||
```
|
||||
|
||||
### Claude Code: `/inbox` skill
|
||||
|
||||
Save as `~/.claude/commands/inbox.md`:
|
||||
|
||||
```markdown
|
||||
---
|
||||
description: Check SynapBus inbox for unread messages. Use at session start.
|
||||
---
|
||||
|
||||
1. Call `my_status` to get unread counts
|
||||
2. If pending DMs exist, read them via execute: `call("read_inbox", {})`
|
||||
3. Summarize what needs attention
|
||||
4. If action items exist, ask user how to proceed
|
||||
```
|
||||
|
||||
### Gemini CLI: Skills
|
||||
|
||||
Save as `~/.gemini/skills/synapbus-check/SKILL.md`:
|
||||
|
||||
```yaml
|
||||
---
|
||||
name: synapbus-check
|
||||
description: Check SynapBus inbox and channel updates
|
||||
---
|
||||
Call my_status to check inbox. Summarize pending DMs and unread channels.
|
||||
If action items exist (priority >= 7), list them.
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Hooks
|
||||
|
||||
### Claude Code: Auto-check inbox on session start
|
||||
|
||||
`.claude/settings.json`:
|
||||
```json
|
||||
{
|
||||
"hooks": {
|
||||
"SessionStart": [
|
||||
{
|
||||
"hooks": [{
|
||||
"type": "command",
|
||||
"command": "echo '{\"hookSpecificOutput\":{\"additionalContext\":\"IMPORTANT: Call my_status on SynapBus MCP to check your inbox before starting work.\"}}'",
|
||||
"timeout": 2000
|
||||
}]
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### Gemini CLI: Session start reminder
|
||||
|
||||
`~/.gemini/settings.json` (add to existing):
|
||||
```json
|
||||
{
|
||||
"hooks": {
|
||||
"SessionStart": [{
|
||||
"hooks": [{
|
||||
"type": "command",
|
||||
"command": "echo '{\"hookSpecificOutput\":{\"additionalContext\":\"Call my_status first to check SynapBus messages.\"}}'",
|
||||
"timeout": 2000
|
||||
}]
|
||||
}]
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Channel Structure
|
||||
|
||||
### Current
|
||||
| Channel | Purpose |
|
||||
|---------|---------|
|
||||
| #general | Cross-cutting discussion |
|
||||
| #open-brain | Long-term memory (509+ entries) |
|
||||
| #approvals | Human approval queue |
|
||||
| #new_posts | Blog post suggestions |
|
||||
| #bugs-synapbus | SynapBus bug reports |
|
||||
| #news-mcpproxy | MCPProxy research |
|
||||
| #news-synapbus | SynapBus research |
|
||||
| #news-personal-brand | Personal brand research |
|
||||
| #reflections-* | Per-agent development reflections |
|
||||
|
||||
### Recommended Additions
|
||||
| Channel | Purpose |
|
||||
|---------|---------|
|
||||
| #bugs-mcpproxy | MCPProxy bug reports |
|
||||
| #bugs-searcher | Searcher pipeline bugs |
|
||||
| #deployments | All deployment announcements |
|
||||
|
||||
---
|
||||
|
||||
## Cross-Agent Communication Pattern
|
||||
|
||||
```
|
||||
Claude Code (dev agent) Gemini CLI (research agent)
|
||||
| |
|
||||
|-- MCP tools ──> SynapBus <── MCP tools --|
|
||||
| (kubic:30088) |
|
||||
| |
|
||||
├─ my_status (check inbox) ├─ my_status |
|
||||
├─ send_message (post/DM) ├─ send_message|
|
||||
├─ search (find context) ├─ search |
|
||||
└─ execute (advanced actions) └─ execute |
|
||||
```
|
||||
|
||||
Both agents connect with their own API keys. SynapBus identifies each by key.
|
||||
Messages, channels, and search are shared — any agent can read any public channel.
|
||||
|
||||
### Example Workflow
|
||||
1. **Gemini research agent** finds a security vulnerability, posts to `#news-mcpproxy`
|
||||
2. **Claude dev agent** starts session, calls `my_status`, sees unread in `#news-mcpproxy`
|
||||
3. Claude reads the finding, assesses impact, fixes the code
|
||||
4. Claude posts fix confirmation to `#news-mcpproxy` as a reply
|
||||
5. Both agents can search for this exchange later via semantic search
|
||||
|
||||
---
|
||||
|
||||
## Protocol Landscape (March 2026)
|
||||
|
||||
| Protocol | Purpose | Relation to SynapBus |
|
||||
|----------|---------|---------------------|
|
||||
| **MCP** | Agent ↔ Tool connectivity | SynapBus IS an MCP server |
|
||||
| **A2A** (Google) | Agent ↔ Agent task delegation | Complementary — A2A for cross-framework; SynapBus for persistent messaging |
|
||||
| **AG-UI** | Agent ↔ Frontend | SynapBus has its own Web UI |
|
||||
| **AGENTS.md** | Agent capability declaration | Could declare SynapBus agents |
|
||||
|
||||
SynapBus sits at the **messaging infrastructure layer**: persistent channels, semantic search, human-observable audit trail. No other MCP server combines all these properties in a single zero-dependency binary.
|
||||
|
||||
---
|
||||
|
||||
## Anti-Patterns
|
||||
|
||||
| Don't | Why |
|
||||
|-------|-----|
|
||||
| Spam channels with progress updates | Floods channels, wastes embedding costs |
|
||||
| Block waiting for agent responses | Other agent may not run for hours |
|
||||
| Send secrets in messages | Messages are stored, searchable, visible in Web UI |
|
||||
| Post same info to multiple channels | Pick the most specific one |
|
||||
| Create channels autonomously | Suggest to human owner instead |
|
||||
| Act on messages > 7 days old without checking for follow-ups | May be already resolved |
|
||||
| Mark everything priority 8+ | Priority inflation kills triage |
|
||||
@@ -0,0 +1,234 @@
|
||||
# SynapBus Roadmap Research — March 2026
|
||||
|
||||
Synthesized findings from 7 parallel research agents covering protocol integration, deployment patterns, enterprise features, and agent coordination.
|
||||
|
||||
---
|
||||
|
||||
## Executive Summary
|
||||
|
||||
| Topic | Key Finding | Priority |
|
||||
|-------|-------------|----------|
|
||||
| **A2A Protocol** | Agent Cards (1-2 days), then inbound gateway (1-2 weeks). Pure Go SDK. | High |
|
||||
| **AG-UI Protocol** | Complement SSE, not replace. Medium-term value. | Low |
|
||||
| **User-level MCP** | Two agents: `claude-algis` + `gemini-algis`. Claude via MCPProxy, Gemini direct. | Do now |
|
||||
| **Mobile access** | Mobile-responsive Web UI via Cloudflare Tunnel. PWA push later. | Medium |
|
||||
| **Cross-device** | Cloudflare Tunnel works for MCP+SSE. Add Cloudflare Access for security. | Do now |
|
||||
| **Always-online agents** | Keep CronJobs + add K8s Job Handlers for reactive response. No daemons. | Medium |
|
||||
| **GitHub Actions** | Only for CI/CD tasks (PR review). K8s is better for research agents. | Low |
|
||||
| **Enterprise IdP** | `coreos/go-oidc/v3` + `golang.org/x/oauth2`. GitHub/Google/Azure AD. | Medium |
|
||||
| **Task acknowledgment** | Claim-process-done for DMs + ACK/DONE convention for channels + StalemateWorker. | High |
|
||||
|
||||
---
|
||||
|
||||
## 1. A2A Protocol Integration
|
||||
|
||||
**What**: Google's Agent-to-Agent protocol (v1.0, 22.6k stars, Linux Foundation).
|
||||
|
||||
**Why**: Makes SynapBus agents discoverable and callable by external frameworks (Google ADK, Microsoft Agent Framework, Strands, LangGraph).
|
||||
|
||||
**Phased approach**:
|
||||
- **Phase 1** (1-2 days): Expose `/.well-known/agent-card.json` from agent registry
|
||||
- **Phase 2** (1-2 weeks): Inbound A2A gateway — external agents send tasks → SynapBus routes as DMs
|
||||
- **Phase 3** (future): Outbound A2A client — SynapBus agents call external A2A agents
|
||||
|
||||
**Key mappings**: A2A Task → SynapBus Conversation, A2A Message → SynapBus Message, A2A Agent Card → SynapBus Agent record.
|
||||
|
||||
**Go SDK**: `github.com/a2aproject/a2a-go` — pure Go, compatible with zero-CGO constraint.
|
||||
|
||||
**vs MCP Tasks (SEP-1686)**: Complementary. MCP Tasks = long-running operations within existing MCP connection. A2A = cross-framework agent interop with discovery.
|
||||
|
||||
---
|
||||
|
||||
## 2. AG-UI Protocol
|
||||
|
||||
**What**: CopilotKit's Agent-User Interaction protocol (12.5k stars). Standardizes agent → frontend streaming.
|
||||
|
||||
**Assessment**: Medium-term value, not urgent. SynapBus's current SSE (notifications) and AG-UI (agent activity streaming) solve different problems.
|
||||
|
||||
**If pursued**: Expose `/ag-ui/run` endpoint that wraps channel activity as AG-UI events. Would let external React frontends (CopilotKit) connect to SynapBus agents.
|
||||
|
||||
**Recommendation**: Watch and plan, but don't build yet. Current SSE + Web UI covers all current use cases.
|
||||
|
||||
---
|
||||
|
||||
## 3. User-Level MCP + Agent Identity
|
||||
|
||||
**Recommendation: Two agent accounts** — `claude-algis` and `gemini-algis`.
|
||||
|
||||
| Tool | SynapBus Access | Agent Identity |
|
||||
|------|----------------|----------------|
|
||||
| Claude Code | Via MCPProxy (user-level, auto-auth) | `claude-algis` |
|
||||
| Gemini CLI | Direct connection (user-level) | `gemini-algis` |
|
||||
| Searcher agents | Direct per-agent keys (unchanged) | `research-*` |
|
||||
|
||||
**Why not one per project**: 20+ projects = 20+ dead agent accounts. **Why not one shared**: Can't tell Claude vs Gemini apart.
|
||||
|
||||
**MCPProxy gateway**: MCPProxy at `localhost:8080` already proxies to kubic. Add `Authorization: Bearer <claude-algis-key>` to the synapbus upstream config in `~/.mcpproxy/mcp_config.json`. All Claude Code projects get SynapBus via BM25 discovery.
|
||||
|
||||
**Gemini**: Direct connection in `~/.gemini/settings.json` with own key.
|
||||
|
||||
**Setup steps**:
|
||||
1. Create agents: `kubectl exec -n synapbus deploy/synapbus -- /synapbus agent create --name claude-algis --display-name "Claude (Algis)" --owner 1`
|
||||
2. Add Bearer header to MCPProxy synapbus upstream
|
||||
3. Remove project-level SynapBus configs from Claude Code
|
||||
4. Add direct SynapBus entry to Gemini settings
|
||||
|
||||
---
|
||||
|
||||
## 4. Mobile Access + Cross-Device
|
||||
|
||||
### Mobile (fastest path)
|
||||
Make Web UI mobile-responsive (sidebar → drawer, touch-friendly compose). Access via `hub.synapbus.dev` on phone. Existing SSE + auth work through Cloudflare Tunnel.
|
||||
|
||||
**Later**: PWA manifest + Web Push for background notifications. iOS supports Web Push since 16.4.
|
||||
|
||||
**Approval on mobile**: Add approve/reject buttons in Web UI for `#approvals` messages (detect `type: "approval_request"` in metadata).
|
||||
|
||||
### Cross-device (home + work)
|
||||
- Home kubic: agents connect locally (`localhost:30088`)
|
||||
- Work laptop: Claude/Gemini connect via `hub.synapbus.dev` tunnel
|
||||
- Benefits: shared context, research feeds dev work, bugs flow between environments
|
||||
|
||||
**Security**: Add Cloudflare Access policy on `hub.synapbus.dev` (email OTP or GitHub SSO). Service tokens for headless agents. OAuth 2.1 remains primary auth layer.
|
||||
|
||||
**Tunnel compatibility**: MCP Streamable HTTP + SSE both work through Cloudflare Tunnel. 30s heartbeats keep connections alive. ~20-50ms round-trip latency.
|
||||
|
||||
---
|
||||
|
||||
## 5. Always-Online Agents
|
||||
|
||||
### Recommended: Hybrid CronJob + K8s Job Handler
|
||||
|
||||
| Workload | Mechanism | Latency | Cost |
|
||||
|----------|-----------|---------|------|
|
||||
| Periodic research sweeps | K8s CronJob (existing) | 4-6h | Low |
|
||||
| Respond to messages/mentions | SynapBus K8s Job Handler | ~10s | Per-event |
|
||||
| Code review/CI tasks | GitHub Actions | ~1m | Free tier |
|
||||
| Always-on daemon | NOT RECOMMENDED | — | High |
|
||||
|
||||
**Keep CronJobs** for scheduled research (already working, staggered schedules).
|
||||
|
||||
**Add K8s Job Handlers** for real-time response: register handlers per agent for `message.received` and `message.mentioned` events. SynapBus spawns K8s Jobs with message context as env vars.
|
||||
|
||||
**Don't use long-running Deployments**: Context windows fill up, resources wasted on single-node MicroK8s.
|
||||
|
||||
**Don't use KEDA**: SynapBus's built-in K8s Job Runner already handles event-driven dispatch.
|
||||
|
||||
### Notable open-source projects
|
||||
- **Kelos**: K8s-native agent orchestration via CRDs (Tasks, AgentConfigs, TaskSpawners)
|
||||
- **Hortator**: Agent reincarnation pattern — checkpoint to `/memory/`, respawn with fresh context
|
||||
- **claude-code-action**: Official GitHub Action for Claude Code in CI/CD
|
||||
|
||||
---
|
||||
|
||||
## 6. Enterprise Identity Providers
|
||||
|
||||
### Architecture
|
||||
```
|
||||
External IdP (GitHub / Google / Azure AD)
|
||||
↓ OIDC Authorization Code Flow
|
||||
SynapBus Identity Layer (NEW: internal/auth/idp/)
|
||||
↓ Creates/links local User + session
|
||||
Existing Auth (Web UI sessions, OAuth AS for MCP, API keys)
|
||||
```
|
||||
|
||||
### Libraries
|
||||
- `coreos/go-oidc/v3` — OIDC discovery + ID token verification (Google, Azure AD)
|
||||
- `golang.org/x/oauth2` — OAuth flow (all providers, already indirect dep)
|
||||
- GitHub: manual OAuth + API calls (not OIDC-compliant)
|
||||
|
||||
### Database
|
||||
```sql
|
||||
CREATE TABLE user_identities (
|
||||
user_id INTEGER REFERENCES users(id),
|
||||
provider TEXT NOT NULL, -- 'github', 'google', 'azuread'
|
||||
external_id TEXT NOT NULL, -- stable provider user ID
|
||||
email TEXT,
|
||||
UNIQUE(provider, external_id)
|
||||
);
|
||||
|
||||
CREATE TABLE identity_providers (
|
||||
id TEXT PRIMARY KEY, -- 'github', 'google', 'azuread-gcore'
|
||||
type TEXT NOT NULL, -- 'github', 'oidc'
|
||||
client_id TEXT NOT NULL,
|
||||
client_secret_encrypted TEXT,
|
||||
issuer_url TEXT, -- OIDC discovery (NULL for GitHub)
|
||||
allowed_domains TEXT, -- '["gcore.com"]'
|
||||
group_mapping TEXT, -- '{"SynapBus-Admins":"admin"}'
|
||||
tenant_id TEXT, -- Azure AD
|
||||
enabled INTEGER DEFAULT 1
|
||||
);
|
||||
```
|
||||
|
||||
### Provider-specific notes
|
||||
- **GitHub**: `read:user` + `user:email` scopes. Map `github_user.id` → external_id.
|
||||
- **Google**: Full OIDC. Restrict to Workspace domain via `hd` claim. Validate server-side.
|
||||
- **Azure AD (Gcore)**: Tenant-specific OIDC. Group claims for role mapping. App Registration in Entra admin center. Handle >200 groups overage.
|
||||
|
||||
### Routes
|
||||
```
|
||||
GET /auth/providers → list enabled IdPs (for login page buttons)
|
||||
GET /auth/login/{provider} → redirect to IdP
|
||||
GET /auth/callback/{provider} → handle callback, create/link user, set session
|
||||
```
|
||||
|
||||
### Multi-tenant: One instance per org (matches local-first philosophy).
|
||||
|
||||
---
|
||||
|
||||
## 7. Task Acknowledgment & Enforcement
|
||||
|
||||
### DM Lifecycle (already built)
|
||||
`pending` → `processing` (claim) → `done` / `failed`
|
||||
|
||||
### CLAUDE.md Instructions (add to all projects)
|
||||
```markdown
|
||||
## Message Acknowledgment (MANDATORY)
|
||||
1. Call `claim_messages` to lock DMs to you
|
||||
2. Process each message
|
||||
3. `mark_done` (success) or `mark_done` with status "failed" + reason
|
||||
4. Never leave claimed messages orphaned — mark failed before session ends
|
||||
```
|
||||
|
||||
### Channel Convention (no code changes)
|
||||
- `ACK: <summary>` — I see it, working on it
|
||||
- `DONE: <summary>` — completed
|
||||
- `BLOCKED: <reason>` — cannot proceed
|
||||
- `DELEGATED: @<agent>` — passed to another agent
|
||||
|
||||
### Enforcement: StalemateWorker (new, small PR)
|
||||
Background worker (like ExpiryWorker/RetentionWorker):
|
||||
- `processing` messages > 24h → auto-fail with "claim timeout"
|
||||
- `pending` messages > 4h → send reminder DM (priority 7)
|
||||
- `pending` messages > 48h → escalate to `#approvals` (priority 9)
|
||||
|
||||
### Channel `reply_to` gap
|
||||
`send_channel_message` action lacks `reply_to` parameter. Add it to enable threaded acknowledgments in channels.
|
||||
|
||||
---
|
||||
|
||||
## Implementation Priority
|
||||
|
||||
### Do Now (zero code)
|
||||
1. Create `claude-algis` + `gemini-algis` agents
|
||||
2. Configure MCPProxy upstream with auth header
|
||||
3. Add acknowledgment protocol to CLAUDE.md / GEMINI.md
|
||||
4. Add SessionStart hooks for inbox checking
|
||||
|
||||
### Next Sprint
|
||||
5. StalemateWorker for message timeout/escalation
|
||||
6. Add `reply_to` to `send_channel_message` action
|
||||
7. A2A Agent Cards (`/.well-known/agent-card.json`)
|
||||
8. Mobile-responsive Web UI (sidebar drawer)
|
||||
|
||||
### Next Month
|
||||
9. A2A inbound gateway (external agents → SynapBus)
|
||||
10. K8s Job Handlers for reactive agent activation
|
||||
11. Enterprise IdP (GitHub + Google + Azure AD)
|
||||
12. PWA with Web Push notifications
|
||||
|
||||
### Future
|
||||
13. A2A outbound client (SynapBus agents → external agents)
|
||||
14. AG-UI endpoint for external frontends
|
||||
15. Telegram bot for mobile approvals
|
||||
16. Approval buttons in Web UI
|
||||
@@ -3,7 +3,11 @@ module github.com/synapbus/synapbus
|
||||
go 1.25.0
|
||||
|
||||
require (
|
||||
github.com/SherClockHolmes/webpush-go v1.4.0
|
||||
github.com/TFMV/hnsw v0.4.0
|
||||
github.com/coreos/go-oidc/v3 v3.17.0
|
||||
github.com/dop251/goja v0.0.0-20260311135729-065cd970411c
|
||||
github.com/evanw/esbuild v0.27.4
|
||||
github.com/go-chi/chi/v5 v5.2.5
|
||||
github.com/google/uuid v1.6.0
|
||||
github.com/mark3labs/mcp-go v0.45.0
|
||||
@@ -12,6 +16,7 @@ require (
|
||||
github.com/prometheus/client_model v0.6.2
|
||||
github.com/spf13/cobra v1.10.2
|
||||
golang.org/x/crypto v0.49.0
|
||||
golang.org/x/oauth2 v0.36.0
|
||||
golang.org/x/time v0.9.0
|
||||
k8s.io/api v0.35.2
|
||||
k8s.io/apimachinery v0.35.2
|
||||
@@ -31,14 +36,13 @@ require (
|
||||
github.com/davecgh/go-spew v1.1.1 // indirect
|
||||
github.com/dgraph-io/ristretto v1.0.0 // indirect
|
||||
github.com/dlclark/regexp2 v1.11.4 // indirect
|
||||
github.com/dop251/goja v0.0.0-20260311135729-065cd970411c // indirect
|
||||
github.com/dustin/go-humanize v1.0.1 // indirect
|
||||
github.com/emicklei/go-restful/v3 v3.12.2 // indirect
|
||||
github.com/evanw/esbuild v0.27.4 // indirect
|
||||
github.com/felixge/httpsnoop v1.0.4 // indirect
|
||||
github.com/fsnotify/fsnotify v1.6.0 // indirect
|
||||
github.com/fxamacker/cbor/v2 v2.9.0 // indirect
|
||||
github.com/go-jose/go-jose/v3 v3.0.3 // indirect
|
||||
github.com/go-jose/go-jose/v4 v4.1.3 // indirect
|
||||
github.com/go-logr/logr v1.4.3 // indirect
|
||||
github.com/go-logr/stdr v1.2.2 // indirect
|
||||
github.com/go-openapi/jsonpointer v0.21.0 // indirect
|
||||
@@ -47,6 +51,7 @@ require (
|
||||
github.com/go-sourcemap/sourcemap v2.1.3+incompatible // indirect
|
||||
github.com/gobuffalo/pop/v6 v6.1.1 // indirect
|
||||
github.com/gogo/protobuf v1.3.2 // indirect
|
||||
github.com/golang-jwt/jwt/v5 v5.2.1 // indirect
|
||||
github.com/golang/mock v1.6.0 // indirect
|
||||
github.com/google/gnostic-models v0.7.0 // indirect
|
||||
github.com/google/pprof v0.0.0-20250403155104-27863c87afa6 // indirect
|
||||
@@ -113,7 +118,6 @@ require (
|
||||
golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546 // indirect
|
||||
golang.org/x/mod v0.33.0 // indirect
|
||||
golang.org/x/net v0.51.0 // indirect
|
||||
golang.org/x/oauth2 v0.30.0 // indirect
|
||||
golang.org/x/sync v0.20.0 // indirect
|
||||
golang.org/x/sys v0.42.0 // indirect
|
||||
golang.org/x/term v0.41.0 // indirect
|
||||
|
||||
@@ -41,6 +41,8 @@ github.com/BurntSushi/xgb v0.0.0-20160522181843-27f122750802/go.mod h1:IVnqGOEym
|
||||
github.com/Masterminds/semver/v3 v3.1.1/go.mod h1:VPu/7SZ7ePZ3QOrcuXROw5FAcLl4a0cBrbBpGY/8hQs=
|
||||
github.com/Masterminds/semver/v3 v3.4.0 h1:Zog+i5UMtVoCU8oKka5P7i9q9HgrJeGzI9SA1Xbatp0=
|
||||
github.com/Masterminds/semver/v3 v3.4.0/go.mod h1:4V+yj/TJE1HU9XfppCwVMZq3I84lprf4nC11bSS5beM=
|
||||
github.com/SherClockHolmes/webpush-go v1.4.0 h1:ocnzNKWN23T9nvHi6IfyrQjkIc0oJWv1B1pULsf9i3s=
|
||||
github.com/SherClockHolmes/webpush-go v1.4.0/go.mod h1:XSq8pKX11vNV8MJEMwjrlTkxhAj1zKfxmyhdV7Pd6UA=
|
||||
github.com/TFMV/hnsw v0.4.0 h1:k61xD3V9LzzwUMDLaHCn+1PbvMbJj33KRdUPiUtuj7k=
|
||||
github.com/TFMV/hnsw v0.4.0/go.mod h1:YPCKBOTpl3KzZxYBTVbR+uH7US5HpprYkDLALt/bgTY=
|
||||
github.com/asaskevich/govalidator v0.0.0-20230301143203-a9d515a09cc2 h1:DklsrG3dyBCFEj5IhUbnKptjxatkF07cF2ak3yi77so=
|
||||
@@ -67,6 +69,8 @@ github.com/cncf/udpa/go v0.0.0-20191209042840-269d4d468f6f/go.mod h1:M8M6+tZqaGX
|
||||
github.com/cncf/udpa/go v0.0.0-20200629203442-efcf912fb354/go.mod h1:WmhPx2Nbnhtbo57+VJT5O0JRkEi1Wbu0z5j0R8u5Hbk=
|
||||
github.com/cncf/udpa/go v0.0.0-20201120205902-5459f2c99403/go.mod h1:WmhPx2Nbnhtbo57+VJT5O0JRkEi1Wbu0z5j0R8u5Hbk=
|
||||
github.com/cockroachdb/apd v1.1.0/go.mod h1:8Sl8LxpKi29FqWXR16WEFZRNSz3SoPzUzeMeY4+DwBQ=
|
||||
github.com/coreos/go-oidc/v3 v3.17.0 h1:hWBGaQfbi0iVviX4ibC7bk8OKT5qNr4klBaCHVNvehc=
|
||||
github.com/coreos/go-oidc/v3 v3.17.0/go.mod h1:wqPbKFrVnE90vty060SB40FCJ8fTHTxSwyXJqZH+sI8=
|
||||
github.com/coreos/go-systemd v0.0.0-20190321100706-95778dfbb74e/go.mod h1:F5haX7vjVVG0kc13fIWeqUViNPyEJxv/OmvnBo0Yme4=
|
||||
github.com/coreos/go-systemd v0.0.0-20190719114852-fd7a80b32e1f/go.mod h1:F5haX7vjVVG0kc13fIWeqUViNPyEJxv/OmvnBo0Yme4=
|
||||
github.com/cpuguy83/go-md2man/v2 v2.0.2/go.mod h1:tgQtvFlXSQOSOSIRvRPT7W67SCa46tRHOmNcaadrF8o=
|
||||
@@ -117,6 +121,8 @@ github.com/go-gl/glfw/v3.3/glfw v0.0.0-20191125211704-12ad95a8df72/go.mod h1:tQ2
|
||||
github.com/go-gl/glfw/v3.3/glfw v0.0.0-20200222043503-6f7a984d4dc4/go.mod h1:tQ2UAYgL5IevRw8kRxooKSPJfGvJ9fJQFa0TUsXzTg8=
|
||||
github.com/go-jose/go-jose/v3 v3.0.3 h1:fFKWeig/irsp7XD2zBxvnmA/XaRWp5V3CBsZXJF7G7k=
|
||||
github.com/go-jose/go-jose/v3 v3.0.3/go.mod h1:5b+7YgP7ZICgJDBdfjZaIt+H/9L9T/YQrVfLAMboGkQ=
|
||||
github.com/go-jose/go-jose/v4 v4.1.3 h1:CVLmWDhDVRa6Mi/IgCgaopNosCaHz7zrMeF9MlZRkrs=
|
||||
github.com/go-jose/go-jose/v4 v4.1.3/go.mod h1:x4oUasVrzR7071A4TnHLGSPpNOm2a21K9Kf04k1rs08=
|
||||
github.com/go-kit/log v0.1.0/go.mod h1:zbhenjAZHb184qTLMA9ZjW7ThYL0H2mk7Q6pNt4vbaY=
|
||||
github.com/go-logfmt/logfmt v0.5.0/go.mod h1:wCYkCAKZfumFQihp8CzCvQ3paCTfi41vtzG1KdI/P7A=
|
||||
github.com/go-logr/logr v1.2.2/go.mod h1:jdQByPbusPIv2/zmleS9BjJVeZ6kBagPoEUsqbVz/1A=
|
||||
@@ -162,6 +168,8 @@ github.com/gofrs/uuid v4.2.0+incompatible/go.mod h1:b2aQJv3Z4Fp6yNu3cdSllBxTCLRx
|
||||
github.com/gofrs/uuid v4.3.1+incompatible/go.mod h1:b2aQJv3Z4Fp6yNu3cdSllBxTCLRxnplIgP/c0N/04lM=
|
||||
github.com/gogo/protobuf v1.3.2 h1:Ov1cvc58UF3b5XjBnZv7+opcTcQFZebYjWzi34vdm4Q=
|
||||
github.com/gogo/protobuf v1.3.2/go.mod h1:P1XiOD3dCwIKUDQYPy72D8LYyHL2YPYrpS2s69NZV8Q=
|
||||
github.com/golang-jwt/jwt/v5 v5.2.1 h1:OuVbFODueb089Lh128TAcimifWaLhJwVflnrgM17wHk=
|
||||
github.com/golang-jwt/jwt/v5 v5.2.1/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk=
|
||||
github.com/golang/glog v0.0.0-20160126235308-23def4e6c14b/go.mod h1:SBH7ygxi8pfUlaOkMMuAQtPIUF8ecWP5IEl/CR7VP2Q=
|
||||
github.com/golang/groupcache v0.0.0-20190702054246-869f871628b6/go.mod h1:cIg4eruTrX1D+g88fzRXU5OdNfaM+9IcxsU14FzY7Hc=
|
||||
github.com/golang/groupcache v0.0.0-20191227052852-215e87163ea7/go.mod h1:cIg4eruTrX1D+g88fzRXU5OdNfaM+9IcxsU14FzY7Hc=
|
||||
@@ -205,6 +213,7 @@ github.com/google/go-cmp v0.5.1/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/
|
||||
github.com/google/go-cmp v0.5.2/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE=
|
||||
github.com/google/go-cmp v0.5.4/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE=
|
||||
github.com/google/go-cmp v0.5.9/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
|
||||
github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
|
||||
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
|
||||
github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
|
||||
github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg=
|
||||
@@ -569,7 +578,10 @@ golang.org/x/crypto v0.0.0-20210616213533-5ff15b29337e/go.mod h1:GvvjBRRGRdwPK5y
|
||||
golang.org/x/crypto v0.0.0-20210711020723-a769d52b0f97/go.mod h1:GvvjBRRGRdwPK5ydBHafDWAxML/pGHZbMvKqRZ5+Abc=
|
||||
golang.org/x/crypto v0.0.0-20210921155107-089bfa567519/go.mod h1:GvvjBRRGRdwPK5ydBHafDWAxML/pGHZbMvKqRZ5+Abc=
|
||||
golang.org/x/crypto v0.0.0-20220722155217-630584e8d5aa/go.mod h1:IxCIyHEi3zRg3s0A5j5BB6A9Jmi73HwBIUl50j+osU4=
|
||||
golang.org/x/crypto v0.13.0/go.mod h1:y6Z2r+Rw4iayiXXAIxJIDAJ1zMW4yaTpebo8fPOliYc=
|
||||
golang.org/x/crypto v0.19.0/go.mod h1:Iy9bg/ha4yyC70EfRS8jz+B6ybOBKMaSxLj6P6oBDfU=
|
||||
golang.org/x/crypto v0.23.0/go.mod h1:CKFgDieR+mRhux2Lsu27y0fO304Db0wZe70UKqHu0v8=
|
||||
golang.org/x/crypto v0.31.0/go.mod h1:kDsLvtWBEx7MV9tJOj9bnXsPbxwJQ6csT/x4KIN4Ssk=
|
||||
golang.org/x/crypto v0.49.0 h1:+Ng2ULVvLHnJ/ZFEq4KdcDd/cfjrrjjNSXNzxg0Y4U4=
|
||||
golang.org/x/crypto v0.49.0/go.mod h1:ErX4dUh2UM+CFYiXZRTcMpEcN8b/1gxEuv3nODoYtCA=
|
||||
golang.org/x/exp v0.0.0-20190121172915-509febef88a4/go.mod h1:CJ0aWSM057203Lf6IL+f9T1iT9GByDxfZKAQTCR3kQA=
|
||||
@@ -611,6 +623,9 @@ golang.org/x/mod v0.4.2/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA=
|
||||
golang.org/x/mod v0.6.0-dev.0.20220419223038-86c51ed26bb4/go.mod h1:jJ57K6gSWd91VN4djpZkiMVwK6gcyfeH4XE8wZrZaV4=
|
||||
golang.org/x/mod v0.8.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
|
||||
golang.org/x/mod v0.10.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
|
||||
golang.org/x/mod v0.12.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
|
||||
golang.org/x/mod v0.15.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c=
|
||||
golang.org/x/mod v0.17.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c=
|
||||
golang.org/x/mod v0.33.0 h1:tHFzIWbBifEmbwtGz65eaWyGiGZatSrT9prnU8DbVL8=
|
||||
golang.org/x/mod v0.33.0/go.mod h1:swjeQEj+6r7fODbD2cqrnje9PnziFuw4bmLbBZFrQ5w=
|
||||
golang.org/x/net v0.0.0-20180724234803-3673e40ba225/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
|
||||
@@ -653,6 +668,9 @@ golang.org/x/net v0.0.0-20221002022538-bcab6841153b/go.mod h1:YDH+HFinaLZZlnHAfS
|
||||
golang.org/x/net v0.6.0/go.mod h1:2Tu9+aMcznHK/AK1HMvgo6xiTLG5rD5rZLDS+rp2Bjs=
|
||||
golang.org/x/net v0.9.0/go.mod h1:d48xBJpPfHeWQsugry2m+kC02ZBRGRgulfHnEXEuWns=
|
||||
golang.org/x/net v0.10.0/go.mod h1:0qNGK6F8kojg2nk9dLZ2mShWaEBan6FAoqfSigmmuDg=
|
||||
golang.org/x/net v0.15.0/go.mod h1:idbUs1IY1+zTqbi8yxTbhexhEEk5ur9LInksu6HrEpk=
|
||||
golang.org/x/net v0.21.0/go.mod h1:bIjVDfnllIU7BJ2DNgfnXvpSvtn8VRwhlsaeUTyUS44=
|
||||
golang.org/x/net v0.25.0/go.mod h1:JkAGAh7GEvH74S6FOH42FLoXpXbE/aqXSrIQjXgsiwM=
|
||||
golang.org/x/net v0.51.0 h1:94R/GTO7mt3/4wIKpcR5gkGmRLOuE/2hNGeWq/GBIFo=
|
||||
golang.org/x/net v0.51.0/go.mod h1:aamm+2QF5ogm02fjy5Bb7CQ0WMt1/WVM7FtyaTLlA9Y=
|
||||
golang.org/x/oauth2 v0.0.0-20180821212333-d2e6202438be/go.mod h1:N/0e6XlmueqKjAGxoOufVs8QHGRruUQn6yWY3a++T0U=
|
||||
@@ -664,8 +682,8 @@ golang.org/x/oauth2 v0.0.0-20200902213428-5d25da1a8d43/go.mod h1:KelEdhl1UZF7XfJ
|
||||
golang.org/x/oauth2 v0.0.0-20201109201403-9fd604954f58/go.mod h1:KelEdhl1UZF7XfJ4dDtk6s++YSgaE7mD/BuKKDLBl4A=
|
||||
golang.org/x/oauth2 v0.0.0-20201208152858-08078c50e5b5/go.mod h1:KelEdhl1UZF7XfJ4dDtk6s++YSgaE7mD/BuKKDLBl4A=
|
||||
golang.org/x/oauth2 v0.0.0-20210218202405-ba52d332ba99/go.mod h1:KelEdhl1UZF7XfJ4dDtk6s++YSgaE7mD/BuKKDLBl4A=
|
||||
golang.org/x/oauth2 v0.30.0 h1:dnDm7JmhM45NNpd8FDDeLhK6FwqbOf4MLCM9zb1BOHI=
|
||||
golang.org/x/oauth2 v0.30.0/go.mod h1:B++QgG3ZKulg6sRPGD/mqlHQs5rB3Ml9erfeDY7xKlU=
|
||||
golang.org/x/oauth2 v0.36.0 h1:peZ/1z27fi9hUOFCAZaHyrpWG5lwe0RJEEEeH0ThlIs=
|
||||
golang.org/x/oauth2 v0.36.0/go.mod h1:YDBUJMTkDnJS+A4BP4eZBjCqtokkg1hODuPjwiGPO7Q=
|
||||
golang.org/x/sync v0.0.0-20180314180146-1d60e4601c6f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.0.0-20181221193216-37e7f081c4d4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
@@ -680,6 +698,10 @@ golang.org/x/sync v0.0.0-20210220032951-036812b2e83c/go.mod h1:RxMgew5VJxzue5/jJ
|
||||
golang.org/x/sync v0.0.0-20220722155255-886fb9371eb4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.0.0-20220929204114-8fcdb60fdcc0/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.1.0/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.3.0/go.mod h1:FU7BRWz2tNW+3quACPkgCx/L+uEAv1htQ0V83Z9Rj+Y=
|
||||
golang.org/x/sync v0.6.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
|
||||
golang.org/x/sync v0.7.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
|
||||
golang.org/x/sync v0.10.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
|
||||
golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4=
|
||||
golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
||||
golang.org/x/sys v0.0.0-20180830151530-49385e6e1522/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||
@@ -736,9 +758,13 @@ golang.org/x/sys v0.5.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.7.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.8.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.12.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.17.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
||||
golang.org/x/sys v0.20.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
||||
golang.org/x/sys v0.28.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
||||
golang.org/x/sys v0.42.0 h1:omrd2nAlyT5ESRdCLYdm3+fMfNFE/+Rf4bDIQImRJeo=
|
||||
golang.org/x/sys v0.42.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
||||
golang.org/x/telemetry v0.0.0-20240228155512-f48c80bd79b2/go.mod h1:TeRTkGYfJXctD9OcfyVLyj2J3IxLnKwHJR8f4D8a3YE=
|
||||
golang.org/x/telemetry v0.0.0-20260209163413-e7419c687ee4 h1:bTLqdHv7xrGlFbvf5/TXNxy/iUwwdkjhqQTJDjW7aj0=
|
||||
golang.org/x/telemetry v0.0.0-20260209163413-e7419c687ee4/go.mod h1:g5NllXBEermZrmR51cJDQxmJUHUOfRAaNyWBM+R+548=
|
||||
golang.org/x/term v0.0.0-20201117132131-f5c789dd3221/go.mod h1:Nr5EML6q2oocZ2LXRh80K7BxOlk5/8JxuGnuhpl+muw=
|
||||
@@ -748,7 +774,10 @@ golang.org/x/term v0.0.0-20220722155259-a9ba230a4035/go.mod h1:jbD1KX2456YbFQfuX
|
||||
golang.org/x/term v0.5.0/go.mod h1:jMB1sMXY+tzblOD4FWmEbocvup2/aLOaQEp7JmGp78k=
|
||||
golang.org/x/term v0.7.0/go.mod h1:P32HKFT3hSsZrRxla30E9HqToFYAQPCMs/zFMBUFqPY=
|
||||
golang.org/x/term v0.8.0/go.mod h1:xPskH00ivmX89bAKVGSKKtLOWNx2+17Eiy94tnKShWo=
|
||||
golang.org/x/term v0.12.0/go.mod h1:owVbMEjm3cBLCHdkQu9b1opXd4ETQWc3BhuQGKgXgvU=
|
||||
golang.org/x/term v0.17.0/go.mod h1:lLRBjIVuehSbZlaOtGMbcMncT+aqLLLmKrsjNrUguwk=
|
||||
golang.org/x/term v0.20.0/go.mod h1:8UkIAJTvZgivsXaD6/pH6U9ecQzZ45awqEOzuCvwpFY=
|
||||
golang.org/x/term v0.27.0/go.mod h1:iMsnZpn0cago0GOrHO2+Y7u7JPn5AylBrcoWkElMTSM=
|
||||
golang.org/x/term v0.41.0 h1:QCgPso/Q3RTJx2Th4bDLqML4W6iJiaXFq2/ftQF13YU=
|
||||
golang.org/x/term v0.41.0/go.mod h1:3pfBgksrReYfZ5lvYM0kSO0LIkAl4Yl2bXOkKP7Ec2A=
|
||||
golang.org/x/text v0.0.0-20170915032832-14c0d48ead0c/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
|
||||
@@ -761,7 +790,10 @@ golang.org/x/text v0.3.6/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
|
||||
golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ=
|
||||
golang.org/x/text v0.7.0/go.mod h1:mrYo+phRRbMaCq/xk9113O4dZlRixOauAjOtrjsXDZ8=
|
||||
golang.org/x/text v0.9.0/go.mod h1:e1OnstbJyHTd6l/uOt8jFFHp6TRDWZR/bV3emEE/zU8=
|
||||
golang.org/x/text v0.13.0/go.mod h1:TvPlkZtksWOMsz7fbANvkp4WM8x/WCo/om8BMLbz+aE=
|
||||
golang.org/x/text v0.14.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU=
|
||||
golang.org/x/text v0.15.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU=
|
||||
golang.org/x/text v0.21.0/go.mod h1:4IBbMaMmOPCJ8SecivzSH54+73PCFmPWxNTLm+vZkEQ=
|
||||
golang.org/x/text v0.35.0 h1:JOVx6vVDFokkpaq1AEptVzLTpDe9KGpj5tR4/X+ybL8=
|
||||
golang.org/x/text v0.35.0/go.mod h1:khi/HExzZJ2pGnjenulevKNX1W67CUy0AsXcNubPGCA=
|
||||
golang.org/x/time v0.0.0-20181108054448-85acf8d2951c/go.mod h1:tRJNPiyCQ0inRvYxbN9jk5I+vvW/OXSQhTDSoE431IQ=
|
||||
@@ -827,6 +859,8 @@ golang.org/x/tools v0.1.1/go.mod h1:o0xws9oXOQQZyjljx8fwUC0k7L1pTE6eaCbjGeHmOkk=
|
||||
golang.org/x/tools v0.1.12/go.mod h1:hNGJHUnrk76NpqgfD5Aqm5Crs+Hm0VOH/i9J2+nxYbc=
|
||||
golang.org/x/tools v0.6.0/go.mod h1:Xwgl3UAJ/d3gWutnCtw505GrjyAbvKui8lOU390QaIU=
|
||||
golang.org/x/tools v0.8.0/go.mod h1:JxBZ99ISMI5ViVkT1tr6tdNmXeTrcpVSD3vZ1RsRdN4=
|
||||
golang.org/x/tools v0.13.0/go.mod h1:HvlwmtVNQAhOuCjW7xxvovg8wbNq7LwfXh/k7wXUl58=
|
||||
golang.org/x/tools v0.21.1-0.20240508182429-e35e4ccd0d2d/go.mod h1:aiJjzUbINMkxbQROHiO6hDPo2LHcIPhhQsa9DLh0yGk=
|
||||
golang.org/x/tools v0.42.0 h1:uNgphsn75Tdz5Ji2q36v/nsFSfR/9BRFvqhGBaJGd5k=
|
||||
golang.org/x/tools v0.42.0/go.mod h1:Ma6lCIwGZvHK6XtgbswSoWroEkhugApmsXyrUmBhfr0=
|
||||
golang.org/x/xerrors v0.0.0-20190410155217-1f06c39b4373/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
||||
@@ -946,6 +980,7 @@ gopkg.in/ini.v1 v1.67.0 h1:Dgnx+6+nfE+IfzjUEISNeydPJh9AXNNsWbGP9KzCsOA=
|
||||
gopkg.in/ini.v1 v1.67.0/go.mod h1:pNLf8WUiyNEtQjuu5G5vTm06TEv9tsIgeAvK8hOrP4k=
|
||||
gopkg.in/yaml.v2 v2.2.2/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI=
|
||||
gopkg.in/yaml.v2 v2.2.4/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI=
|
||||
gopkg.in/yaml.v2 v2.4.0 h1:D8xgwECY7CYvx+Y2n4sBz93Jn9JRvxdiyyo8CTfuKaY=
|
||||
gopkg.in/yaml.v2 v2.4.0/go.mod h1:RDklbk79AGWmwhnvt/jBztapEOGDOx6ZbXqjP6csGnQ=
|
||||
gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
||||
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
|
||||
|
||||
@@ -0,0 +1,133 @@
|
||||
// Package a2a provides A2A Agent Card discovery for SynapBus.
|
||||
// The Agent Card is a JSON document that describes the hub and its registered
|
||||
// agents, following the A2A Agent Card specification.
|
||||
package a2a
|
||||
|
||||
import "encoding/json"
|
||||
|
||||
// AgentCard is the A2A Agent Card document returned by the discovery endpoint.
|
||||
type AgentCard struct {
|
||||
Name string `json:"name"`
|
||||
Description string `json:"description"`
|
||||
Version string `json:"version"`
|
||||
SupportedInterfaces []AgentInterface `json:"supported_interfaces"`
|
||||
Capabilities AgentCapabilities `json:"capabilities"`
|
||||
Skills []AgentSkill `json:"skills"`
|
||||
SecuritySchemes map[string]any `json:"security_schemes"`
|
||||
DefaultInputModes []string `json:"default_input_modes"`
|
||||
DefaultOutputModes []string `json:"default_output_modes"`
|
||||
}
|
||||
|
||||
// AgentInterface describes a protocol endpoint the hub supports.
|
||||
type AgentInterface struct {
|
||||
URL string `json:"url"`
|
||||
ProtocolBinding string `json:"protocol_binding"`
|
||||
}
|
||||
|
||||
// AgentCapabilities declares hub-level capabilities.
|
||||
type AgentCapabilities struct {
|
||||
Streaming bool `json:"streaming"`
|
||||
PushNotifications bool `json:"push_notifications"`
|
||||
}
|
||||
|
||||
// AgentSkill represents a single agent registered on the hub, mapped as an
|
||||
// A2A skill.
|
||||
type AgentSkill struct {
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Description string `json:"description,omitempty"`
|
||||
Tags []string `json:"tags,omitempty"`
|
||||
}
|
||||
|
||||
// AgentInfo is a lightweight struct used to pass agent data from the registry
|
||||
// into the card generator without leaking internal types.
|
||||
type AgentInfo struct {
|
||||
Name string
|
||||
DisplayName string
|
||||
Type string
|
||||
Capabilities json.RawMessage
|
||||
}
|
||||
|
||||
// GenerateAgentCard builds an AgentCard from the hub configuration and a list
|
||||
// of registered agents.
|
||||
func GenerateAgentCard(baseURL string, version string, agents []AgentInfo) *AgentCard {
|
||||
skills := make([]AgentSkill, 0, len(agents))
|
||||
for _, a := range agents {
|
||||
skill := AgentSkill{
|
||||
ID: a.Name,
|
||||
Name: a.DisplayName,
|
||||
}
|
||||
if skill.Name == "" {
|
||||
skill.Name = a.Name
|
||||
}
|
||||
|
||||
// Parse capabilities JSON for description and tags.
|
||||
if len(a.Capabilities) > 0 {
|
||||
var caps map[string]interface{}
|
||||
if json.Unmarshal(a.Capabilities, &caps) == nil {
|
||||
if desc, ok := caps["description"].(string); ok {
|
||||
skill.Description = desc
|
||||
}
|
||||
if role, ok := caps["role"].(string); ok {
|
||||
skill.Tags = append(skill.Tags, role)
|
||||
}
|
||||
if tagsRaw, ok := caps["tags"]; ok {
|
||||
switch v := tagsRaw.(type) {
|
||||
case []interface{}:
|
||||
for _, t := range v {
|
||||
if s, ok := t.(string); ok {
|
||||
skill.Tags = append(skill.Tags, s)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Add agent type as a tag.
|
||||
if a.Type != "" {
|
||||
skill.Tags = append(skill.Tags, a.Type)
|
||||
}
|
||||
|
||||
skills = append(skills, skill)
|
||||
}
|
||||
|
||||
return &AgentCard{
|
||||
Name: "SynapBus Hub",
|
||||
Description: "MCP-native agent-to-agent messaging hub",
|
||||
Version: version,
|
||||
SupportedInterfaces: []AgentInterface{
|
||||
{
|
||||
URL: baseURL + "/a2a",
|
||||
ProtocolBinding: "JSONRPC",
|
||||
},
|
||||
},
|
||||
Capabilities: AgentCapabilities{
|
||||
Streaming: true,
|
||||
PushNotifications: false,
|
||||
},
|
||||
Skills: skills,
|
||||
SecuritySchemes: map[string]any{
|
||||
"apiKey": map[string]any{
|
||||
"type": "apiKey",
|
||||
"in": "header",
|
||||
"name": "Authorization",
|
||||
"scheme": "Bearer",
|
||||
},
|
||||
"oauth2": map[string]any{
|
||||
"type": "oauth2",
|
||||
"flows": map[string]any{
|
||||
"authorizationCode": map[string]any{
|
||||
"authorizationUrl": baseURL + "/oauth/authorize",
|
||||
"tokenUrl": baseURL + "/oauth/token",
|
||||
"scopes": map[string]string{
|
||||
"mcp": "MCP protocol access",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
DefaultInputModes: []string{"text"},
|
||||
DefaultOutputModes: []string{"text"},
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,218 @@
|
||||
package a2a
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestGenerateAgentCard_WithAgents(t *testing.T) {
|
||||
agents := []AgentInfo{
|
||||
{Name: "research-bot", DisplayName: "Research Bot", Type: "ai", Capabilities: json.RawMessage(`{"role":"researcher","description":"Searches the web"}`)},
|
||||
{Name: "social-commenter", DisplayName: "Social Commenter", Type: "ai", Capabilities: json.RawMessage(`{"tags":["social","marketing"]}`)},
|
||||
{Name: "data-analyst", DisplayName: "", Type: "ai", Capabilities: json.RawMessage(`{}`)},
|
||||
}
|
||||
|
||||
card := GenerateAgentCard("http://localhost:8080", "1.0.0", agents)
|
||||
|
||||
if card.Name != "SynapBus Hub" {
|
||||
t.Errorf("name = %q, want %q", card.Name, "SynapBus Hub")
|
||||
}
|
||||
if card.Version != "1.0.0" {
|
||||
t.Errorf("version = %q, want %q", card.Version, "1.0.0")
|
||||
}
|
||||
if len(card.Skills) != 3 {
|
||||
t.Fatalf("skills count = %d, want 3", len(card.Skills))
|
||||
}
|
||||
|
||||
// Verify first skill has description and tags from capabilities
|
||||
s0 := card.Skills[0]
|
||||
if s0.ID != "research-bot" {
|
||||
t.Errorf("skill[0].id = %q, want %q", s0.ID, "research-bot")
|
||||
}
|
||||
if s0.Name != "Research Bot" {
|
||||
t.Errorf("skill[0].name = %q, want %q", s0.Name, "Research Bot")
|
||||
}
|
||||
if s0.Description != "Searches the web" {
|
||||
t.Errorf("skill[0].description = %q, want %q", s0.Description, "Searches the web")
|
||||
}
|
||||
// Should have "researcher" from role + "ai" from type
|
||||
if len(s0.Tags) < 2 {
|
||||
t.Errorf("skill[0].tags = %v, expected at least 2 tags", s0.Tags)
|
||||
}
|
||||
|
||||
// Second skill should have tags from capabilities "tags" field
|
||||
s1 := card.Skills[1]
|
||||
if s1.ID != "social-commenter" {
|
||||
t.Errorf("skill[1].id = %q, want %q", s1.ID, "social-commenter")
|
||||
}
|
||||
foundSocial := false
|
||||
for _, tag := range s1.Tags {
|
||||
if tag == "social" {
|
||||
foundSocial = true
|
||||
}
|
||||
}
|
||||
if !foundSocial {
|
||||
t.Errorf("skill[1].tags = %v, expected 'social' tag", s1.Tags)
|
||||
}
|
||||
|
||||
// Third skill should use name as display name (since DisplayName is empty)
|
||||
s2 := card.Skills[2]
|
||||
if s2.Name != "data-analyst" {
|
||||
t.Errorf("skill[2].name = %q, want %q (fallback to Name)", s2.Name, "data-analyst")
|
||||
}
|
||||
|
||||
// Verify JSON serialization round-trips cleanly
|
||||
data, err := json.Marshal(card)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal card: %v", err)
|
||||
}
|
||||
var decoded AgentCard
|
||||
if err := json.Unmarshal(data, &decoded); err != nil {
|
||||
t.Fatalf("unmarshal card: %v", err)
|
||||
}
|
||||
if decoded.Name != card.Name {
|
||||
t.Errorf("round-trip name mismatch")
|
||||
}
|
||||
if len(decoded.Skills) != 3 {
|
||||
t.Errorf("round-trip skills count = %d, want 3", len(decoded.Skills))
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerateAgentCard_WithCapabilities_TagsPopulated(t *testing.T) {
|
||||
agents := []AgentInfo{
|
||||
{
|
||||
Name: "smart-agent",
|
||||
DisplayName: "Smart Agent",
|
||||
Type: "ai",
|
||||
Capabilities: json.RawMessage(`{"role":"analyst","description":"Analyzes data","tags":["ml","data"]}`),
|
||||
},
|
||||
}
|
||||
|
||||
card := GenerateAgentCard("http://example.com", "2.0.0", agents)
|
||||
|
||||
if len(card.Skills) != 1 {
|
||||
t.Fatalf("skills count = %d, want 1", len(card.Skills))
|
||||
}
|
||||
|
||||
skill := card.Skills[0]
|
||||
if skill.Description != "Analyzes data" {
|
||||
t.Errorf("description = %q, want %q", skill.Description, "Analyzes data")
|
||||
}
|
||||
|
||||
// Expect tags: "analyst" (from role), "ml", "data" (from tags), "ai" (from type)
|
||||
expectedTags := map[string]bool{"analyst": false, "ml": false, "data": false, "ai": false}
|
||||
for _, tag := range skill.Tags {
|
||||
if _, ok := expectedTags[tag]; ok {
|
||||
expectedTags[tag] = true
|
||||
}
|
||||
}
|
||||
for tag, found := range expectedTags {
|
||||
if !found {
|
||||
t.Errorf("missing expected tag %q in %v", tag, skill.Tags)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerateAgentCard_NoAgents(t *testing.T) {
|
||||
card := GenerateAgentCard("http://localhost:8080", "0.1.0", nil)
|
||||
|
||||
if card.Name != "SynapBus Hub" {
|
||||
t.Errorf("name = %q, want %q", card.Name, "SynapBus Hub")
|
||||
}
|
||||
if len(card.Skills) != 0 {
|
||||
t.Errorf("skills count = %d, want 0", len(card.Skills))
|
||||
}
|
||||
if len(card.SupportedInterfaces) != 1 {
|
||||
t.Fatalf("interfaces count = %d, want 1", len(card.SupportedInterfaces))
|
||||
}
|
||||
if card.SupportedInterfaces[0].URL != "http://localhost:8080/a2a" {
|
||||
t.Errorf("interface url = %q, want %q", card.SupportedInterfaces[0].URL, "http://localhost:8080/a2a")
|
||||
}
|
||||
if card.SecuritySchemes == nil {
|
||||
t.Error("security_schemes should not be nil")
|
||||
}
|
||||
}
|
||||
|
||||
// mockAgentLister implements AgentLister for handler tests.
|
||||
type mockAgentLister struct {
|
||||
agents []AgentInfo
|
||||
err error
|
||||
}
|
||||
|
||||
func (m *mockAgentLister) ListAllActiveAgents(_ context.Context) ([]AgentInfo, error) {
|
||||
return m.agents, m.err
|
||||
}
|
||||
|
||||
func TestHandler_Returns200WithCorrectContentType(t *testing.T) {
|
||||
lister := &mockAgentLister{
|
||||
agents: []AgentInfo{
|
||||
{Name: "bot-1", DisplayName: "Bot One", Type: "ai"},
|
||||
{Name: "bot-2", DisplayName: "Bot Two", Type: "ai"},
|
||||
},
|
||||
}
|
||||
|
||||
handler := NewAgentCardHandler(lister, "http://localhost:8080", "1.0.0")
|
||||
req := httptest.NewRequest(http.MethodGet, "/.well-known/agent-card.json", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
handler.ServeHTTP(rr, req)
|
||||
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Errorf("status = %d, want %d", rr.Code, http.StatusOK)
|
||||
}
|
||||
ct := rr.Header().Get("Content-Type")
|
||||
if ct != "application/json" {
|
||||
t.Errorf("Content-Type = %q, want %q", ct, "application/json")
|
||||
}
|
||||
|
||||
var card AgentCard
|
||||
if err := json.NewDecoder(rr.Body).Decode(&card); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if card.Name != "SynapBus Hub" {
|
||||
t.Errorf("card.name = %q, want %q", card.Name, "SynapBus Hub")
|
||||
}
|
||||
if len(card.Skills) != 2 {
|
||||
t.Errorf("card.skills count = %d, want 2", len(card.Skills))
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandler_DerivesBaseURLFromRequest(t *testing.T) {
|
||||
lister := &mockAgentLister{agents: nil}
|
||||
|
||||
// Empty configuredBaseURL — should derive from request
|
||||
handler := NewAgentCardHandler(lister, "", "1.0.0")
|
||||
req := httptest.NewRequest(http.MethodGet, "/.well-known/agent-card.json", nil)
|
||||
req.Host = "myhost:9090"
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
handler.ServeHTTP(rr, req)
|
||||
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want %d", rr.Code, http.StatusOK)
|
||||
}
|
||||
|
||||
var card AgentCard
|
||||
if err := json.NewDecoder(rr.Body).Decode(&card); err != nil {
|
||||
t.Fatalf("decode: %v", err)
|
||||
}
|
||||
if card.SupportedInterfaces[0].URL != "http://myhost:9090/a2a" {
|
||||
t.Errorf("interface url = %q, want %q", card.SupportedInterfaces[0].URL, "http://myhost:9090/a2a")
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandler_RejectsNonGET(t *testing.T) {
|
||||
lister := &mockAgentLister{agents: nil}
|
||||
handler := NewAgentCardHandler(lister, "http://localhost:8080", "1.0.0")
|
||||
req := httptest.NewRequest(http.MethodPost, "/.well-known/agent-card.json", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
handler.ServeHTTP(rr, req)
|
||||
|
||||
if rr.Code != http.StatusMethodNotAllowed {
|
||||
t.Errorf("status = %d, want %d", rr.Code, http.StatusMethodNotAllowed)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,306 @@
|
||||
package a2a
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
|
||||
"github.com/google/uuid"
|
||||
|
||||
"github.com/synapbus/synapbus/internal/agents"
|
||||
"github.com/synapbus/synapbus/internal/messaging"
|
||||
)
|
||||
|
||||
// MessagingService defines the messaging operations needed by the A2A gateway.
|
||||
type MessagingService interface {
|
||||
SendMessage(ctx context.Context, from, to, body string, opts messaging.SendOptions) (*messaging.Message, error)
|
||||
GetConversation(ctx context.Context, id int64) (*messaging.Conversation, []*messaging.Message, error)
|
||||
}
|
||||
|
||||
// AgentService defines the agent operations needed by the A2A gateway.
|
||||
type AgentService interface {
|
||||
GetAgent(ctx context.Context, name string) (*agents.Agent, error)
|
||||
}
|
||||
|
||||
// Gateway handles inbound A2A JSON-RPC requests.
|
||||
type Gateway struct {
|
||||
taskStore *A2ATaskStore
|
||||
msgService MessagingService
|
||||
agentService AgentService
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
// NewGateway creates a new A2A gateway.
|
||||
func NewGateway(taskStore *A2ATaskStore, msgService MessagingService, agentService AgentService) *Gateway {
|
||||
return &Gateway{
|
||||
taskStore: taskStore,
|
||||
msgService: msgService,
|
||||
agentService: agentService,
|
||||
logger: slog.Default().With("component", "a2a-gateway"),
|
||||
}
|
||||
}
|
||||
|
||||
// JSON-RPC 2.0 request/response types.
|
||||
|
||||
type jsonRPCRequest struct {
|
||||
JSONRPC string `json:"jsonrpc"`
|
||||
ID json.RawMessage `json:"id"`
|
||||
Method string `json:"method"`
|
||||
Params json.RawMessage `json:"params,omitempty"`
|
||||
}
|
||||
|
||||
type jsonRPCResponse struct {
|
||||
JSONRPC string `json:"jsonrpc"`
|
||||
ID json.RawMessage `json:"id"`
|
||||
Result any `json:"result,omitempty"`
|
||||
Error *jsonRPCError `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
type jsonRPCError struct {
|
||||
Code int `json:"code"`
|
||||
Message string `json:"message"`
|
||||
}
|
||||
|
||||
// Standard JSON-RPC 2.0 error codes.
|
||||
const (
|
||||
errCodeParse = -32700
|
||||
errCodeInvalidReq = -32600
|
||||
errCodeNoMethod = -32601
|
||||
errCodeInvalidParams = -32602
|
||||
errCodeInternal = -32603
|
||||
)
|
||||
|
||||
// HandleJSONRPC dispatches incoming JSON-RPC 2.0 requests to the appropriate handler.
|
||||
func (g *Gateway) HandleJSONRPC(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
http.Error(w, `{"error":"method not allowed"}`, http.StatusMethodNotAllowed)
|
||||
return
|
||||
}
|
||||
|
||||
body, err := io.ReadAll(io.LimitReader(r.Body, 1<<20)) // 1 MiB limit
|
||||
if err != nil {
|
||||
writeJSONRPC(w, nil, nil, &jsonRPCError{Code: errCodeParse, Message: "failed to read request body"})
|
||||
return
|
||||
}
|
||||
|
||||
var req jsonRPCRequest
|
||||
if err := json.Unmarshal(body, &req); err != nil {
|
||||
writeJSONRPC(w, nil, nil, &jsonRPCError{Code: errCodeParse, Message: "invalid JSON"})
|
||||
return
|
||||
}
|
||||
|
||||
if req.JSONRPC != "2.0" {
|
||||
writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInvalidReq, Message: "jsonrpc must be \"2.0\""})
|
||||
return
|
||||
}
|
||||
|
||||
// Extract the calling agent from auth context.
|
||||
callerAgent, ok := agents.AgentFromContext(r.Context())
|
||||
callerName := ""
|
||||
if ok && callerAgent != nil {
|
||||
callerName = callerAgent.Name
|
||||
}
|
||||
|
||||
switch req.Method {
|
||||
case "message.send":
|
||||
g.handleMessageSend(w, r.Context(), req, callerName)
|
||||
case "tasks.get":
|
||||
g.handleTasksGet(w, r.Context(), req)
|
||||
case "tasks.cancel":
|
||||
g.handleTasksCancel(w, r.Context(), req)
|
||||
default:
|
||||
writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeNoMethod, Message: fmt.Sprintf("unknown method: %s", req.Method)})
|
||||
}
|
||||
}
|
||||
|
||||
// message.send params
|
||||
|
||||
type messageSendParams struct {
|
||||
Message struct {
|
||||
Body string `json:"body"`
|
||||
Metadata struct {
|
||||
TargetAgent string `json:"target_agent"`
|
||||
} `json:"metadata"`
|
||||
} `json:"message"`
|
||||
}
|
||||
|
||||
func (g *Gateway) handleMessageSend(w http.ResponseWriter, ctx context.Context, req jsonRPCRequest, callerName string) {
|
||||
var params messageSendParams
|
||||
if err := json.Unmarshal(req.Params, ¶ms); err != nil {
|
||||
writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInvalidParams, Message: "invalid params: " + err.Error()})
|
||||
return
|
||||
}
|
||||
|
||||
targetAgent := params.Message.Metadata.TargetAgent
|
||||
if targetAgent == "" {
|
||||
writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInvalidParams, Message: "params.message.metadata.target_agent is required"})
|
||||
return
|
||||
}
|
||||
|
||||
messageBody := params.Message.Body
|
||||
if messageBody == "" {
|
||||
writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInvalidParams, Message: "params.message.body is required"})
|
||||
return
|
||||
}
|
||||
|
||||
// Validate target agent exists.
|
||||
_, err := g.agentService.GetAgent(ctx, targetAgent)
|
||||
if err != nil {
|
||||
writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInvalidParams, Message: fmt.Sprintf("target agent not found: %s", targetAgent)})
|
||||
return
|
||||
}
|
||||
|
||||
// Create task.
|
||||
taskID := uuid.New().String()
|
||||
contextID := uuid.New().String()
|
||||
|
||||
// Determine the sender name for the DM. Use the caller's identity if
|
||||
// authenticated, otherwise fall back to "a2a-gateway" so SendMessage
|
||||
// has a non-empty from field.
|
||||
senderName := callerName
|
||||
if senderName == "" {
|
||||
senderName = "a2a-gateway"
|
||||
}
|
||||
|
||||
// Build metadata containing the a2a_task_id.
|
||||
metaJSON, _ := json.Marshal(map[string]string{"a2a_task_id": taskID})
|
||||
|
||||
// Send DM to target agent.
|
||||
msg, err := g.msgService.SendMessage(ctx, senderName, targetAgent, messageBody, messaging.SendOptions{
|
||||
Subject: fmt.Sprintf("A2A Task %s", taskID),
|
||||
Metadata: string(metaJSON),
|
||||
})
|
||||
if err != nil {
|
||||
g.logger.Error("failed to send DM for A2A task", "task_id", taskID, "error", err)
|
||||
writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInternal, Message: "failed to deliver message: " + err.Error()})
|
||||
return
|
||||
}
|
||||
|
||||
convID := msg.ConversationID
|
||||
task := &A2ATask{
|
||||
ID: taskID,
|
||||
ContextID: contextID,
|
||||
TargetAgent: targetAgent,
|
||||
SourceAgent: callerName,
|
||||
ConversationID: &convID,
|
||||
State: StateSubmitted,
|
||||
}
|
||||
|
||||
if err := g.taskStore.CreateTask(ctx, task); err != nil {
|
||||
g.logger.Error("failed to create A2A task", "task_id", taskID, "error", err)
|
||||
writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInternal, Message: "failed to create task"})
|
||||
return
|
||||
}
|
||||
|
||||
g.logger.Info("A2A task created",
|
||||
"task_id", taskID,
|
||||
"target_agent", targetAgent,
|
||||
"source_agent", callerName,
|
||||
"message_id", msg.ID,
|
||||
)
|
||||
|
||||
writeJSONRPC(w, req.ID, task, nil)
|
||||
}
|
||||
|
||||
// tasks.get params
|
||||
|
||||
type tasksGetParams struct {
|
||||
ID string `json:"id"`
|
||||
}
|
||||
|
||||
func (g *Gateway) handleTasksGet(w http.ResponseWriter, ctx context.Context, req jsonRPCRequest) {
|
||||
var params tasksGetParams
|
||||
if err := json.Unmarshal(req.Params, ¶ms); err != nil {
|
||||
writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInvalidParams, Message: "invalid params: " + err.Error()})
|
||||
return
|
||||
}
|
||||
|
||||
if params.ID == "" {
|
||||
writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInvalidParams, Message: "params.id is required"})
|
||||
return
|
||||
}
|
||||
|
||||
task, err := g.taskStore.GetTask(ctx, params.ID)
|
||||
if err != nil {
|
||||
writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInvalidParams, Message: err.Error()})
|
||||
return
|
||||
}
|
||||
|
||||
// If the task has a conversation and is still SUBMITTED, check whether the
|
||||
// target agent has replied, which indicates completion.
|
||||
if task.State == StateSubmitted && task.ConversationID != nil {
|
||||
_, msgs, err := g.msgService.GetConversation(ctx, *task.ConversationID)
|
||||
if err == nil && len(msgs) > 1 {
|
||||
// Check if the target agent sent a reply (any message from target after the first).
|
||||
for _, m := range msgs[1:] {
|
||||
if m.FromAgent == task.TargetAgent {
|
||||
task.State = StateCompleted
|
||||
_ = g.taskStore.UpdateTaskState(ctx, task.ID, StateCompleted)
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
writeJSONRPC(w, req.ID, task, nil)
|
||||
}
|
||||
|
||||
// tasks.cancel params
|
||||
|
||||
type tasksCancelParams struct {
|
||||
ID string `json:"id"`
|
||||
}
|
||||
|
||||
func (g *Gateway) handleTasksCancel(w http.ResponseWriter, ctx context.Context, req jsonRPCRequest) {
|
||||
var params tasksCancelParams
|
||||
if err := json.Unmarshal(req.Params, ¶ms); err != nil {
|
||||
writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInvalidParams, Message: "invalid params: " + err.Error()})
|
||||
return
|
||||
}
|
||||
|
||||
if params.ID == "" {
|
||||
writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInvalidParams, Message: "params.id is required"})
|
||||
return
|
||||
}
|
||||
|
||||
task, err := g.taskStore.GetTask(ctx, params.ID)
|
||||
if err != nil {
|
||||
writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInvalidParams, Message: err.Error()})
|
||||
return
|
||||
}
|
||||
|
||||
// Cannot cancel a terminal task.
|
||||
if task.State == StateCompleted || task.State == StateCanceled {
|
||||
writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInvalidParams, Message: fmt.Sprintf("task is already in terminal state: %s", task.State)})
|
||||
return
|
||||
}
|
||||
|
||||
if err := g.taskStore.UpdateTaskState(ctx, task.ID, StateCanceled); err != nil {
|
||||
writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInternal, Message: "failed to cancel task"})
|
||||
return
|
||||
}
|
||||
|
||||
task.State = StateCanceled
|
||||
g.logger.Info("A2A task canceled", "task_id", task.ID)
|
||||
|
||||
writeJSONRPC(w, req.ID, task, nil)
|
||||
}
|
||||
|
||||
// writeJSONRPC writes a JSON-RPC 2.0 response.
|
||||
func writeJSONRPC(w http.ResponseWriter, id json.RawMessage, result any, rpcErr *jsonRPCError) {
|
||||
resp := jsonRPCResponse{
|
||||
JSONRPC: "2.0",
|
||||
ID: id,
|
||||
Result: result,
|
||||
Error: rpcErr,
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
if rpcErr != nil {
|
||||
// Use 200 for JSON-RPC errors (per spec), but set result to nil.
|
||||
resp.Result = nil
|
||||
}
|
||||
json.NewEncoder(w).Encode(resp)
|
||||
}
|
||||
@@ -0,0 +1,483 @@
|
||||
package a2a
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
|
||||
"github.com/synapbus/synapbus/internal/agents"
|
||||
"github.com/synapbus/synapbus/internal/messaging"
|
||||
"github.com/synapbus/synapbus/internal/storage"
|
||||
)
|
||||
|
||||
// --- test helpers ---
|
||||
|
||||
func newTestDB(t *testing.T) *sql.DB {
|
||||
t.Helper()
|
||||
dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", t.Name())
|
||||
db, err := sql.Open("sqlite", dsn)
|
||||
if err != nil {
|
||||
t.Fatalf("open database: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { db.Close() })
|
||||
|
||||
if _, err := db.Exec("PRAGMA foreign_keys=ON"); err != nil {
|
||||
t.Fatalf("enable foreign keys: %v", err)
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
if err := storage.RunMigrations(ctx, db); err != nil {
|
||||
t.Fatalf("run migrations: %v", err)
|
||||
}
|
||||
|
||||
// Seed a test user for owner_id FK
|
||||
db.Exec(`INSERT OR IGNORE INTO users (id, username, password_hash, display_name) VALUES (1, 'testowner', 'hash', 'Test Owner')`)
|
||||
|
||||
return db
|
||||
}
|
||||
|
||||
func seedAgent(t *testing.T, db *sql.DB, name string) {
|
||||
t.Helper()
|
||||
_, err := db.Exec(
|
||||
`INSERT OR IGNORE INTO agents (name, display_name, type, capabilities, owner_id, api_key_hash, status) VALUES (?, ?, 'ai', '{}', 1, 'testhash', 'active')`,
|
||||
name, name,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("seed agent %s: %v", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
// mockMsgService implements MessagingService for testing.
|
||||
type mockMsgService struct {
|
||||
lastFrom string
|
||||
lastTo string
|
||||
lastBody string
|
||||
lastOpts messaging.SendOptions
|
||||
sendErr error
|
||||
returnMsg *messaging.Message
|
||||
convMsgs []*messaging.Message
|
||||
getConvErr error
|
||||
}
|
||||
|
||||
func (m *mockMsgService) SendMessage(_ context.Context, from, to, body string, opts messaging.SendOptions) (*messaging.Message, error) {
|
||||
m.lastFrom = from
|
||||
m.lastTo = to
|
||||
m.lastBody = body
|
||||
m.lastOpts = opts
|
||||
if m.sendErr != nil {
|
||||
return nil, m.sendErr
|
||||
}
|
||||
if m.returnMsg != nil {
|
||||
return m.returnMsg, nil
|
||||
}
|
||||
return &messaging.Message{
|
||||
ID: 1,
|
||||
ConversationID: 100,
|
||||
FromAgent: from,
|
||||
ToAgent: to,
|
||||
Body: body,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (m *mockMsgService) GetConversation(_ context.Context, id int64) (*messaging.Conversation, []*messaging.Message, error) {
|
||||
if m.getConvErr != nil {
|
||||
return nil, nil, m.getConvErr
|
||||
}
|
||||
conv := &messaging.Conversation{ID: id}
|
||||
return conv, m.convMsgs, nil
|
||||
}
|
||||
|
||||
// mockAgentService implements AgentService for testing.
|
||||
type mockAgentService struct {
|
||||
agents map[string]*agents.Agent
|
||||
}
|
||||
|
||||
func (m *mockAgentService) GetAgent(_ context.Context, name string) (*agents.Agent, error) {
|
||||
if a, ok := m.agents[name]; ok {
|
||||
return a, nil
|
||||
}
|
||||
return nil, fmt.Errorf("agent not found: %s", name)
|
||||
}
|
||||
|
||||
func newTestGateway(t *testing.T) (*Gateway, *mockMsgService, *mockAgentService, *A2ATaskStore) {
|
||||
t.Helper()
|
||||
db := newTestDB(t)
|
||||
seedAgent(t, db, "target-bot")
|
||||
seedAgent(t, db, "sender-bot")
|
||||
|
||||
taskStore := NewA2ATaskStore(db)
|
||||
msgSvc := &mockMsgService{}
|
||||
agentSvc := &mockAgentService{
|
||||
agents: map[string]*agents.Agent{
|
||||
"target-bot": {ID: 1, Name: "target-bot", Status: "active"},
|
||||
"sender-bot": {ID: 2, Name: "sender-bot", Status: "active"},
|
||||
},
|
||||
}
|
||||
|
||||
gw := NewGateway(taskStore, msgSvc, agentSvc)
|
||||
return gw, msgSvc, agentSvc, taskStore
|
||||
}
|
||||
|
||||
func jsonRPCCall(method string, params any) []byte {
|
||||
p, _ := json.Marshal(params)
|
||||
req := map[string]any{
|
||||
"jsonrpc": "2.0",
|
||||
"id": 1,
|
||||
"method": method,
|
||||
"params": json.RawMessage(p),
|
||||
}
|
||||
b, _ := json.Marshal(req)
|
||||
return b
|
||||
}
|
||||
|
||||
func doRequest(t *testing.T, gw *Gateway, body []byte, agentCtx *agents.Agent) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
req := httptest.NewRequest(http.MethodPost, "/a2a", bytes.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
if agentCtx != nil {
|
||||
req = req.WithContext(agents.ContextWithAgent(req.Context(), agentCtx))
|
||||
}
|
||||
w := httptest.NewRecorder()
|
||||
gw.HandleJSONRPC(w, req)
|
||||
return w
|
||||
}
|
||||
|
||||
func parseResponse(t *testing.T, w *httptest.ResponseRecorder) jsonRPCResponse {
|
||||
t.Helper()
|
||||
var resp jsonRPCResponse
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
|
||||
t.Fatalf("unmarshal response: %v\nbody: %s", err, w.Body.String())
|
||||
}
|
||||
return resp
|
||||
}
|
||||
|
||||
// --- tests ---
|
||||
|
||||
func TestMessageSend_CreatesTaskAndDM(t *testing.T) {
|
||||
gw, msgSvc, _, taskStore := newTestGateway(t)
|
||||
|
||||
body := jsonRPCCall("message.send", map[string]any{
|
||||
"message": map[string]any{
|
||||
"body": "Hello target bot",
|
||||
"metadata": map[string]string{
|
||||
"target_agent": "target-bot",
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
callerAgent := &agents.Agent{Name: "sender-bot"}
|
||||
w := doRequest(t, gw, body, callerAgent)
|
||||
resp := parseResponse(t, w)
|
||||
|
||||
if resp.Error != nil {
|
||||
t.Fatalf("unexpected error: code=%d message=%s", resp.Error.Code, resp.Error.Message)
|
||||
}
|
||||
|
||||
// Verify a task was returned.
|
||||
resultBytes, _ := json.Marshal(resp.Result)
|
||||
var task A2ATask
|
||||
if err := json.Unmarshal(resultBytes, &task); err != nil {
|
||||
t.Fatalf("unmarshal task result: %v", err)
|
||||
}
|
||||
|
||||
if task.ID == "" {
|
||||
t.Error("task ID should not be empty")
|
||||
}
|
||||
if task.State != StateSubmitted {
|
||||
t.Errorf("task state = %q, want %q", task.State, StateSubmitted)
|
||||
}
|
||||
if task.TargetAgent != "target-bot" {
|
||||
t.Errorf("target_agent = %q, want %q", task.TargetAgent, "target-bot")
|
||||
}
|
||||
if task.SourceAgent != "sender-bot" {
|
||||
t.Errorf("source_agent = %q, want %q", task.SourceAgent, "sender-bot")
|
||||
}
|
||||
|
||||
// Verify the DM was sent.
|
||||
if msgSvc.lastTo != "target-bot" {
|
||||
t.Errorf("DM to = %q, want %q", msgSvc.lastTo, "target-bot")
|
||||
}
|
||||
if msgSvc.lastFrom != "sender-bot" {
|
||||
t.Errorf("DM from = %q, want %q", msgSvc.lastFrom, "sender-bot")
|
||||
}
|
||||
if msgSvc.lastBody != "Hello target bot" {
|
||||
t.Errorf("DM body = %q, want %q", msgSvc.lastBody, "Hello target bot")
|
||||
}
|
||||
|
||||
// Verify metadata contains a2a_task_id.
|
||||
var meta map[string]string
|
||||
if err := json.Unmarshal([]byte(msgSvc.lastOpts.Metadata), &meta); err != nil {
|
||||
t.Fatalf("unmarshal metadata: %v", err)
|
||||
}
|
||||
if meta["a2a_task_id"] != task.ID {
|
||||
t.Errorf("metadata a2a_task_id = %q, want %q", meta["a2a_task_id"], task.ID)
|
||||
}
|
||||
|
||||
// Verify task is persisted.
|
||||
stored, err := taskStore.GetTask(context.Background(), task.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetTask: %v", err)
|
||||
}
|
||||
if stored.State != StateSubmitted {
|
||||
t.Errorf("stored task state = %q, want %q", stored.State, StateSubmitted)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTasksGet_ReturnsTask(t *testing.T) {
|
||||
gw, _, _, taskStore := newTestGateway(t)
|
||||
|
||||
// Create a task directly.
|
||||
convID := int64(100)
|
||||
task := &A2ATask{
|
||||
ID: "test-task-123",
|
||||
ContextID: "ctx-123",
|
||||
TargetAgent: "target-bot",
|
||||
SourceAgent: "sender-bot",
|
||||
ConversationID: &convID,
|
||||
State: StateSubmitted,
|
||||
}
|
||||
if err := taskStore.CreateTask(context.Background(), task); err != nil {
|
||||
t.Fatalf("CreateTask: %v", err)
|
||||
}
|
||||
|
||||
body := jsonRPCCall("tasks.get", map[string]string{"id": "test-task-123"})
|
||||
w := doRequest(t, gw, body, nil)
|
||||
resp := parseResponse(t, w)
|
||||
|
||||
if resp.Error != nil {
|
||||
t.Fatalf("unexpected error: code=%d message=%s", resp.Error.Code, resp.Error.Message)
|
||||
}
|
||||
|
||||
resultBytes, _ := json.Marshal(resp.Result)
|
||||
var got A2ATask
|
||||
if err := json.Unmarshal(resultBytes, &got); err != nil {
|
||||
t.Fatalf("unmarshal task result: %v", err)
|
||||
}
|
||||
|
||||
if got.ID != "test-task-123" {
|
||||
t.Errorf("task ID = %q, want %q", got.ID, "test-task-123")
|
||||
}
|
||||
if got.State != StateSubmitted {
|
||||
t.Errorf("task state = %q, want %q", got.State, StateSubmitted)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTasksGet_CompletesOnReply(t *testing.T) {
|
||||
gw, msgSvc, _, taskStore := newTestGateway(t)
|
||||
|
||||
// Create a task.
|
||||
convID := int64(100)
|
||||
task := &A2ATask{
|
||||
ID: "task-reply-test",
|
||||
ContextID: "ctx-456",
|
||||
TargetAgent: "target-bot",
|
||||
SourceAgent: "sender-bot",
|
||||
ConversationID: &convID,
|
||||
State: StateSubmitted,
|
||||
}
|
||||
if err := taskStore.CreateTask(context.Background(), task); err != nil {
|
||||
t.Fatalf("CreateTask: %v", err)
|
||||
}
|
||||
|
||||
// Simulate the target agent having replied.
|
||||
msgSvc.convMsgs = []*messaging.Message{
|
||||
{ID: 1, FromAgent: "sender-bot", Body: "Hello"},
|
||||
{ID: 2, FromAgent: "target-bot", Body: "Reply from target"},
|
||||
}
|
||||
|
||||
body := jsonRPCCall("tasks.get", map[string]string{"id": "task-reply-test"})
|
||||
w := doRequest(t, gw, body, nil)
|
||||
resp := parseResponse(t, w)
|
||||
|
||||
if resp.Error != nil {
|
||||
t.Fatalf("unexpected error: code=%d message=%s", resp.Error.Code, resp.Error.Message)
|
||||
}
|
||||
|
||||
resultBytes, _ := json.Marshal(resp.Result)
|
||||
var got A2ATask
|
||||
if err := json.Unmarshal(resultBytes, &got); err != nil {
|
||||
t.Fatalf("unmarshal: %v", err)
|
||||
}
|
||||
|
||||
if got.State != StateCompleted {
|
||||
t.Errorf("task state = %q, want %q (target agent replied)", got.State, StateCompleted)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTasksCancel_TransitionsToCanceled(t *testing.T) {
|
||||
gw, _, _, taskStore := newTestGateway(t)
|
||||
|
||||
task := &A2ATask{
|
||||
ID: "task-cancel-test",
|
||||
ContextID: "ctx-789",
|
||||
TargetAgent: "target-bot",
|
||||
SourceAgent: "sender-bot",
|
||||
State: StateSubmitted,
|
||||
}
|
||||
if err := taskStore.CreateTask(context.Background(), task); err != nil {
|
||||
t.Fatalf("CreateTask: %v", err)
|
||||
}
|
||||
|
||||
body := jsonRPCCall("tasks.cancel", map[string]string{"id": "task-cancel-test"})
|
||||
w := doRequest(t, gw, body, nil)
|
||||
resp := parseResponse(t, w)
|
||||
|
||||
if resp.Error != nil {
|
||||
t.Fatalf("unexpected error: code=%d message=%s", resp.Error.Code, resp.Error.Message)
|
||||
}
|
||||
|
||||
resultBytes, _ := json.Marshal(resp.Result)
|
||||
var got A2ATask
|
||||
if err := json.Unmarshal(resultBytes, &got); err != nil {
|
||||
t.Fatalf("unmarshal: %v", err)
|
||||
}
|
||||
|
||||
if got.State != StateCanceled {
|
||||
t.Errorf("task state = %q, want %q", got.State, StateCanceled)
|
||||
}
|
||||
|
||||
// Verify persisted state.
|
||||
stored, err := taskStore.GetTask(context.Background(), "task-cancel-test")
|
||||
if err != nil {
|
||||
t.Fatalf("GetTask: %v", err)
|
||||
}
|
||||
if stored.State != StateCanceled {
|
||||
t.Errorf("stored state = %q, want %q", stored.State, StateCanceled)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTasksCancel_TerminalStateError(t *testing.T) {
|
||||
gw, _, _, taskStore := newTestGateway(t)
|
||||
|
||||
task := &A2ATask{
|
||||
ID: "task-already-done",
|
||||
ContextID: "ctx-done",
|
||||
TargetAgent: "target-bot",
|
||||
State: StateCompleted,
|
||||
}
|
||||
if err := taskStore.CreateTask(context.Background(), task); err != nil {
|
||||
t.Fatalf("CreateTask: %v", err)
|
||||
}
|
||||
|
||||
body := jsonRPCCall("tasks.cancel", map[string]string{"id": "task-already-done"})
|
||||
w := doRequest(t, gw, body, nil)
|
||||
resp := parseResponse(t, w)
|
||||
|
||||
if resp.Error == nil {
|
||||
t.Fatal("expected error for canceling terminal task")
|
||||
}
|
||||
if resp.Error.Code != errCodeInvalidParams {
|
||||
t.Errorf("error code = %d, want %d", resp.Error.Code, errCodeInvalidParams)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMessageSend_NonExistentAgent(t *testing.T) {
|
||||
gw, _, _, _ := newTestGateway(t)
|
||||
|
||||
body := jsonRPCCall("message.send", map[string]any{
|
||||
"message": map[string]any{
|
||||
"body": "Hello ghost",
|
||||
"metadata": map[string]string{
|
||||
"target_agent": "does-not-exist",
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
w := doRequest(t, gw, body, nil)
|
||||
resp := parseResponse(t, w)
|
||||
|
||||
if resp.Error == nil {
|
||||
t.Fatal("expected error for non-existent target agent")
|
||||
}
|
||||
if resp.Error.Code != errCodeInvalidParams {
|
||||
t.Errorf("error code = %d, want %d", resp.Error.Code, errCodeInvalidParams)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInvalidJSONRPC(t *testing.T) {
|
||||
gw, _, _, _ := newTestGateway(t)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
body string
|
||||
}{
|
||||
{
|
||||
name: "not JSON",
|
||||
body: "this is not json",
|
||||
},
|
||||
{
|
||||
name: "wrong jsonrpc version",
|
||||
body: `{"jsonrpc":"1.0","id":1,"method":"message.send","params":{}}`,
|
||||
},
|
||||
{
|
||||
name: "unknown method",
|
||||
body: `{"jsonrpc":"2.0","id":1,"method":"unknown.method","params":{}}`,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodPost, "/a2a", bytes.NewBufferString(tt.body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
gw.HandleJSONRPC(w, req)
|
||||
|
||||
var resp jsonRPCResponse
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
|
||||
t.Fatalf("unmarshal response: %v\nbody: %s", err, w.Body.String())
|
||||
}
|
||||
if resp.Error == nil {
|
||||
t.Error("expected error response")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestUnauthenticatedRequest_NoAgentContext(t *testing.T) {
|
||||
// This tests that message.send works even without an authenticated agent
|
||||
// in context (caller is anonymous), using "a2a-gateway" as the sender.
|
||||
gw, msgSvc, _, _ := newTestGateway(t)
|
||||
|
||||
body := jsonRPCCall("message.send", map[string]any{
|
||||
"message": map[string]any{
|
||||
"body": "Hello from anonymous",
|
||||
"metadata": map[string]string{
|
||||
"target_agent": "target-bot",
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
// No agent context — simulates unauthenticated-at-gateway-level
|
||||
// (in practice, the auth middleware would block this; this tests the
|
||||
// gateway's fallback behavior).
|
||||
w := doRequest(t, gw, body, nil)
|
||||
resp := parseResponse(t, w)
|
||||
|
||||
if resp.Error != nil {
|
||||
t.Fatalf("unexpected error: code=%d message=%s", resp.Error.Code, resp.Error.Message)
|
||||
}
|
||||
|
||||
// Should use "a2a-gateway" as sender when no caller agent.
|
||||
if msgSvc.lastFrom != "a2a-gateway" {
|
||||
t.Errorf("DM from = %q, want %q", msgSvc.lastFrom, "a2a-gateway")
|
||||
}
|
||||
}
|
||||
|
||||
func TestHTTPMethodNotAllowed(t *testing.T) {
|
||||
gw, _, _, _ := newTestGateway(t)
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/a2a", nil)
|
||||
w := httptest.NewRecorder()
|
||||
gw.HandleJSONRPC(w, req)
|
||||
|
||||
if w.Code != http.StatusMethodNotAllowed {
|
||||
t.Errorf("status = %d, want %d", w.Code, http.StatusMethodNotAllowed)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,56 @@
|
||||
package a2a
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
)
|
||||
|
||||
// AgentLister abstracts the operation of listing active non-human agents.
|
||||
type AgentLister interface {
|
||||
ListAllActiveAgents(ctx context.Context) ([]AgentInfo, error)
|
||||
}
|
||||
|
||||
// NewAgentCardHandler returns an http.HandlerFunc that serves the A2A Agent
|
||||
// Card JSON document at /.well-known/agent-card.json.
|
||||
//
|
||||
// The handler is public (no auth required) because Agent Cards are meant for
|
||||
// discovery. If configuredBaseURL is empty the base URL is derived from the
|
||||
// incoming request (respecting X-Forwarded-* headers).
|
||||
func NewAgentCardHandler(agentLister AgentLister, configuredBaseURL string, version string) http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet {
|
||||
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
|
||||
return
|
||||
}
|
||||
|
||||
// Derive base URL from request if not configured.
|
||||
baseURL := configuredBaseURL
|
||||
if baseURL == "" {
|
||||
scheme := "http"
|
||||
if r.TLS != nil {
|
||||
scheme = "https"
|
||||
}
|
||||
if proto := r.Header.Get("X-Forwarded-Proto"); proto != "" {
|
||||
scheme = proto
|
||||
}
|
||||
host := r.Host
|
||||
if fwdHost := r.Header.Get("X-Forwarded-Host"); fwdHost != "" {
|
||||
host = fwdHost
|
||||
}
|
||||
baseURL = scheme + "://" + host
|
||||
}
|
||||
|
||||
agents, err := agentLister.ListAllActiveAgents(r.Context())
|
||||
if err != nil {
|
||||
http.Error(w, "internal error", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
card := GenerateAgentCard(baseURL, version, agents)
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.Header().Set("Cache-Control", "public, max-age=60")
|
||||
json.NewEncoder(w).Encode(card)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,132 @@
|
||||
// Package a2a provides the A2A (Agent-to-Agent) inbound gateway for SynapBus.
|
||||
// External A2A-compliant agents can send tasks to SynapBus agents via JSON-RPC.
|
||||
package a2a
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Task states following the A2A protocol.
|
||||
const (
|
||||
StateSubmitted = "SUBMITTED"
|
||||
StateCompleted = "COMPLETED"
|
||||
StateCanceled = "CANCELED"
|
||||
)
|
||||
|
||||
// A2ATask represents an inbound A2A task tracked by the gateway.
|
||||
type A2ATask struct {
|
||||
ID string `json:"id"`
|
||||
ContextID string `json:"context_id"`
|
||||
TargetAgent string `json:"target_agent"`
|
||||
SourceAgent string `json:"source_agent"`
|
||||
ConversationID *int64 `json:"conversation_id,omitempty"`
|
||||
State string `json:"state"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
// A2ATaskStore provides CRUD operations for A2A tasks backed by SQLite.
|
||||
type A2ATaskStore struct {
|
||||
db *sql.DB
|
||||
}
|
||||
|
||||
// NewA2ATaskStore creates a new task store.
|
||||
func NewA2ATaskStore(db *sql.DB) *A2ATaskStore {
|
||||
return &A2ATaskStore{db: db}
|
||||
}
|
||||
|
||||
// CreateTask inserts a new A2A task into the database.
|
||||
func (s *A2ATaskStore) CreateTask(ctx context.Context, task *A2ATask) error {
|
||||
_, err := s.db.ExecContext(ctx,
|
||||
`INSERT INTO a2a_tasks (id, context_id, target_agent, source_agent, conversation_id, state, created_at, updated_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`,
|
||||
task.ID, task.ContextID, task.TargetAgent, task.SourceAgent, task.ConversationID, task.State,
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("insert a2a task: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetTask returns an A2A task by its ID.
|
||||
func (s *A2ATaskStore) GetTask(ctx context.Context, id string) (*A2ATask, error) {
|
||||
var task A2ATask
|
||||
var conversationID sql.NullInt64
|
||||
err := s.db.QueryRowContext(ctx,
|
||||
`SELECT id, context_id, target_agent, source_agent, conversation_id, state, created_at, updated_at
|
||||
FROM a2a_tasks WHERE id = ?`, id,
|
||||
).Scan(&task.ID, &task.ContextID, &task.TargetAgent, &task.SourceAgent,
|
||||
&conversationID, &task.State, &task.CreatedAt, &task.UpdatedAt)
|
||||
if err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, fmt.Errorf("a2a task not found: %s", id)
|
||||
}
|
||||
return nil, fmt.Errorf("get a2a task: %w", err)
|
||||
}
|
||||
if conversationID.Valid {
|
||||
task.ConversationID = &conversationID.Int64
|
||||
}
|
||||
return &task, nil
|
||||
}
|
||||
|
||||
// UpdateTaskState transitions a task to a new state.
|
||||
func (s *A2ATaskStore) UpdateTaskState(ctx context.Context, id, state string) error {
|
||||
result, err := s.db.ExecContext(ctx,
|
||||
`UPDATE a2a_tasks SET state = ?, updated_at = CURRENT_TIMESTAMP WHERE id = ?`,
|
||||
state, id,
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("update a2a task state: %w", err)
|
||||
}
|
||||
rows, err := result.RowsAffected()
|
||||
if err != nil {
|
||||
return fmt.Errorf("get rows affected: %w", err)
|
||||
}
|
||||
if rows == 0 {
|
||||
return fmt.Errorf("a2a task not found: %s", id)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ListTasks returns tasks for a target agent, optionally filtered by state.
|
||||
func (s *A2ATaskStore) ListTasks(ctx context.Context, targetAgent, state string) ([]*A2ATask, error) {
|
||||
var query string
|
||||
var args []any
|
||||
|
||||
if state != "" {
|
||||
query = `SELECT id, context_id, target_agent, source_agent, conversation_id, state, created_at, updated_at
|
||||
FROM a2a_tasks WHERE target_agent = ? AND state = ? ORDER BY created_at DESC`
|
||||
args = []any{targetAgent, state}
|
||||
} else {
|
||||
query = `SELECT id, context_id, target_agent, source_agent, conversation_id, state, created_at, updated_at
|
||||
FROM a2a_tasks WHERE target_agent = ? ORDER BY created_at DESC`
|
||||
args = []any{targetAgent}
|
||||
}
|
||||
|
||||
rows, err := s.db.QueryContext(ctx, query, args...)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list a2a tasks: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var tasks []*A2ATask
|
||||
for rows.Next() {
|
||||
var task A2ATask
|
||||
var conversationID sql.NullInt64
|
||||
if err := rows.Scan(&task.ID, &task.ContextID, &task.TargetAgent, &task.SourceAgent,
|
||||
&conversationID, &task.State, &task.CreatedAt, &task.UpdatedAt); err != nil {
|
||||
return nil, fmt.Errorf("scan a2a task: %w", err)
|
||||
}
|
||||
if conversationID.Valid {
|
||||
task.ConversationID = &conversationID.Int64
|
||||
}
|
||||
tasks = append(tasks, &task)
|
||||
}
|
||||
if tasks == nil {
|
||||
tasks = []*A2ATask{}
|
||||
}
|
||||
return tasks, rows.Err()
|
||||
}
|
||||
@@ -6,10 +6,10 @@ type Registry struct {
|
||||
ordered []Action // maintains insertion order
|
||||
}
|
||||
|
||||
// NewRegistry creates a registry pre-populated with all 23 agent-callable actions.
|
||||
// NewRegistry creates a registry pre-populated with all 27 agent-callable actions.
|
||||
func NewRegistry() *Registry {
|
||||
r := &Registry{
|
||||
actions: make(map[string]Action, 23),
|
||||
actions: make(map[string]Action, 27),
|
||||
}
|
||||
for _, a := range allActions() {
|
||||
r.actions[a.Name] = a
|
||||
@@ -42,7 +42,7 @@ func (r *Registry) ListByCategory(category string) []Action {
|
||||
return out
|
||||
}
|
||||
|
||||
// allActions returns the canonical list of all 23 agent-callable actions.
|
||||
// allActions returns the canonical list of all 27 agent-callable actions.
|
||||
func allActions() []Action {
|
||||
return []Action{
|
||||
// ── Messaging (7 actions) ──────────────────────────────────────
|
||||
@@ -306,6 +306,7 @@ func allActions() []Action {
|
||||
{Name: "body", Type: "string", Description: "Message body text", Required: true},
|
||||
{Name: "priority", Type: "number", Description: "Message priority (1-10, default 5)", Default: "5"},
|
||||
{Name: "metadata", Type: "string", Description: "JSON metadata object (optional)"},
|
||||
{Name: "reply_to", Type: "number", Description: "Message ID to reply to (creates a thread)", Required: false},
|
||||
},
|
||||
Returns: "JSON with channel_id, message_id, and status 'sent'",
|
||||
Examples: []Example{
|
||||
@@ -424,7 +425,7 @@ func allActions() []Action {
|
||||
{
|
||||
Name: "upload_attachment",
|
||||
Category: "attachments",
|
||||
Description: "Upload a file attachment. Content must be base64-encoded. Returns the SHA-256 hash for later retrieval. Max file size: 50MB.",
|
||||
Description: "Upload a file attachment. Content must be base64-encoded. Returns the SHA-256 hash for later retrieval. Upload first, then use the returned hash in send_message's attachments parameter to link it to a message. Max file size: 50MB.",
|
||||
Params: []Param{
|
||||
{Name: "content", Type: "string", Description: "Base64-encoded file content", Required: true},
|
||||
{Name: "filename", Type: "string", Description: "Original filename (optional, used for MIME detection and display)"},
|
||||
@@ -454,5 +455,75 @@ func allActions() []Action {
|
||||
},
|
||||
},
|
||||
},
|
||||
|
||||
// ── Reactions (4 actions) ────────────────────────────────────
|
||||
{
|
||||
Name: "react",
|
||||
Category: "reactions",
|
||||
Description: "Add or toggle a reaction on a message. Valid reactions: approve, reject, in_progress, done, published. Adding the same reaction again removes it (toggle).",
|
||||
Params: []Param{
|
||||
{Name: "message_id", Type: "number", Description: "ID of the message to react to", Required: true},
|
||||
{Name: "reaction", Type: "string", Description: "Reaction type: approve, reject, in_progress, done, published", Required: true},
|
||||
{Name: "metadata", Type: "string", Description: "JSON metadata object (optional)"},
|
||||
},
|
||||
Returns: "JSON with action ('added' or 'removed') and reaction details",
|
||||
Examples: []Example{
|
||||
{
|
||||
Description: "Approve a message",
|
||||
Code: `call("react", {"message_id": 42, "reaction": "approve"})`,
|
||||
},
|
||||
{
|
||||
Description: "Toggle a reaction off (call same reaction again)",
|
||||
Code: `call("react", {"message_id": 42, "reaction": "approve"})`,
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "unreact",
|
||||
Category: "reactions",
|
||||
Description: "Remove a specific reaction from a message.",
|
||||
Params: []Param{
|
||||
{Name: "message_id", Type: "number", Description: "ID of the message to remove reaction from", Required: true},
|
||||
{Name: "reaction", Type: "string", Description: "Reaction type to remove: approve, reject, in_progress, done, published", Required: true},
|
||||
},
|
||||
Returns: "JSON with message_id, reaction, and status 'removed'",
|
||||
Examples: []Example{
|
||||
{
|
||||
Description: "Remove an approval reaction",
|
||||
Code: `call("unreact", {"message_id": 42, "reaction": "approve"})`,
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "get_reactions",
|
||||
Category: "reactions",
|
||||
Description: "Get all reactions on a message and its derived workflow state.",
|
||||
Params: []Param{
|
||||
{Name: "message_id", Type: "number", Description: "ID of the message to get reactions for", Required: true},
|
||||
},
|
||||
Returns: "JSON with reactions array and workflow_state",
|
||||
Examples: []Example{
|
||||
{
|
||||
Description: "Get reactions and workflow state for a message",
|
||||
Code: `call("get_reactions", {"message_id": 42})`,
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "list_by_state",
|
||||
Category: "reactions",
|
||||
Description: "List messages in a channel filtered by workflow state. Valid states: proposed, approved, in_progress, rejected, done, published.",
|
||||
Params: []Param{
|
||||
{Name: "channel", Type: "string", Description: "Channel name", Required: true},
|
||||
{Name: "state", Type: "string", Description: "Workflow state to filter by: proposed, approved, in_progress, rejected, done, published", Required: true},
|
||||
},
|
||||
Returns: "JSON with message_ids array and count",
|
||||
Examples: []Example{
|
||||
{
|
||||
Description: "List approved messages in a channel",
|
||||
Code: `call("list_by_state", {"channel": "approvals", "state": "approved"})`,
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
@@ -4,11 +4,11 @@ import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestRegistryHas23Actions(t *testing.T) {
|
||||
func TestRegistryHas27Actions(t *testing.T) {
|
||||
r := NewRegistry()
|
||||
got := len(r.List())
|
||||
if got != 23 {
|
||||
t.Errorf("expected 23 actions, got %d", got)
|
||||
if got != 27 {
|
||||
t.Errorf("expected 27 actions, got %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -23,6 +23,7 @@ func TestRegistryCategories(t *testing.T) {
|
||||
{"channels", 9},
|
||||
{"swarm", 5},
|
||||
{"attachments", 2},
|
||||
{"reactions", 4},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
@@ -49,6 +50,8 @@ func TestRegistryGetByName(t *testing.T) {
|
||||
"post_task", "bid_task", "accept_bid", "complete_task", "list_tasks",
|
||||
// attachments
|
||||
"upload_attachment", "download_attachment",
|
||||
// reactions
|
||||
"react", "unreact", "get_reactions", "list_by_state",
|
||||
}
|
||||
|
||||
for _, name := range allNames {
|
||||
|
||||
@@ -139,6 +139,8 @@ func (s *AdminServer) dispatch(req Request) Response {
|
||||
return s.handleAgentDelete(ctx, req.Args)
|
||||
case "agent.revoke_key":
|
||||
return s.handleAgentRevokeKey(ctx, req.Args)
|
||||
case "agent.update_capabilities":
|
||||
return s.handleAgentUpdateCapabilities(ctx, req.Args)
|
||||
|
||||
// --- audit commands ---
|
||||
case "audit.list":
|
||||
@@ -167,6 +169,8 @@ func (s *AdminServer) dispatch(req Request) Response {
|
||||
return s.handleChannelsCreate(ctx, req.Args)
|
||||
case "channels.join":
|
||||
return s.handleChannelsJoin(ctx, req.Args)
|
||||
case "channels.update_settings":
|
||||
return s.handleChannelsUpdateSettings(ctx, req.Args)
|
||||
|
||||
// --- conversations ---
|
||||
case "conversations.list":
|
||||
@@ -446,6 +450,35 @@ func (s *AdminServer) handleAgentRevokeKey(ctx context.Context, args json.RawMes
|
||||
}}
|
||||
}
|
||||
|
||||
func (s *AdminServer) handleAgentUpdateCapabilities(ctx context.Context, args json.RawMessage) Response {
|
||||
var p struct {
|
||||
Name string `json:"name"`
|
||||
Capabilities json.RawMessage `json:"capabilities"`
|
||||
}
|
||||
if err := json.Unmarshal(args, &p); err != nil {
|
||||
return Response{OK: false, Error: "invalid args: " + err.Error()}
|
||||
}
|
||||
if p.Name == "" {
|
||||
return Response{OK: false, Error: "name is required"}
|
||||
}
|
||||
if len(p.Capabilities) == 0 {
|
||||
return Response{OK: false, Error: "capabilities is required"}
|
||||
}
|
||||
if !json.Valid(p.Capabilities) {
|
||||
return Response{OK: false, Error: "capabilities must be valid JSON"}
|
||||
}
|
||||
|
||||
agent, err := s.services.Agents.UpdateAgent(ctx, p.Name, "", p.Capabilities)
|
||||
if err != nil {
|
||||
return Response{OK: false, Error: err.Error()}
|
||||
}
|
||||
|
||||
return Response{OK: true, Data: map[string]interface{}{
|
||||
"name": agent.Name,
|
||||
"capabilities": json.RawMessage(agent.Capabilities),
|
||||
}}
|
||||
}
|
||||
|
||||
// ---------- audit handlers ----------
|
||||
|
||||
func (s *AdminServer) handleAuditList(ctx context.Context, args json.RawMessage) Response {
|
||||
@@ -964,6 +997,64 @@ func (s *AdminServer) handleChannelsJoin(ctx context.Context, args json.RawMessa
|
||||
}}
|
||||
}
|
||||
|
||||
func (s *AdminServer) handleChannelsUpdateSettings(ctx context.Context, args json.RawMessage) Response {
|
||||
var p struct {
|
||||
Name string `json:"name"`
|
||||
AutoApprove *bool `json:"auto_approve,omitempty"`
|
||||
StalemateRemindAfter string `json:"stalemate_remind_after,omitempty"`
|
||||
StalemateEscalateAfter string `json:"stalemate_escalate_after,omitempty"`
|
||||
}
|
||||
if err := json.Unmarshal(args, &p); err != nil {
|
||||
return Response{OK: false, Error: "invalid args: " + err.Error()}
|
||||
}
|
||||
if p.Name == "" {
|
||||
return Response{OK: false, Error: "name is required"}
|
||||
}
|
||||
|
||||
// Build the SET clause dynamically based on provided fields
|
||||
var setClauses []string
|
||||
var setArgs []interface{}
|
||||
|
||||
if p.AutoApprove != nil {
|
||||
autoApproveVal := 0
|
||||
if *p.AutoApprove {
|
||||
autoApproveVal = 1
|
||||
}
|
||||
setClauses = append(setClauses, "auto_approve = ?")
|
||||
setArgs = append(setArgs, autoApproveVal)
|
||||
}
|
||||
if p.StalemateRemindAfter != "" {
|
||||
setClauses = append(setClauses, "stalemate_remind_after = ?")
|
||||
setArgs = append(setArgs, p.StalemateRemindAfter)
|
||||
}
|
||||
if p.StalemateEscalateAfter != "" {
|
||||
setClauses = append(setClauses, "stalemate_escalate_after = ?")
|
||||
setArgs = append(setArgs, p.StalemateEscalateAfter)
|
||||
}
|
||||
|
||||
if len(setClauses) == 0 {
|
||||
return Response{OK: false, Error: "at least one setting must be provided (auto_approve, stalemate_remind_after, stalemate_escalate_after)"}
|
||||
}
|
||||
|
||||
query := fmt.Sprintf("UPDATE channels SET %s WHERE LOWER(name) = LOWER(?)", strings.Join(setClauses, ", "))
|
||||
setArgs = append(setArgs, p.Name)
|
||||
|
||||
result, err := s.db.ExecContext(ctx, query, setArgs...)
|
||||
if err != nil {
|
||||
return Response{OK: false, Error: "update channel settings: " + err.Error()}
|
||||
}
|
||||
|
||||
rowsAffected, _ := result.RowsAffected()
|
||||
if rowsAffected == 0 {
|
||||
return Response{OK: false, Error: fmt.Sprintf("channel not found: %s", p.Name)}
|
||||
}
|
||||
|
||||
return Response{OK: true, Data: map[string]interface{}{
|
||||
"channel": p.Name,
|
||||
"updated": true,
|
||||
}}
|
||||
}
|
||||
|
||||
// ---------- conversations handlers ----------
|
||||
|
||||
func (s *AdminServer) handleConversationsList(ctx context.Context, args json.RawMessage) Response {
|
||||
|
||||
@@ -229,9 +229,13 @@ func RequiredAuthMiddlewareWithOAuth(service *AgentService, keyService *apikeys.
|
||||
|
||||
// resolveOAuthToken introspects an OAuth bearer token and extracts agent identity.
|
||||
func resolveOAuthToken(ctx context.Context, provider fosite.OAuth2Provider, token string, service *AgentService) (agentName string, ownerID string, ok bool) {
|
||||
// Use a fositeSession-compatible struct for introspection.
|
||||
// We import the type indirectly through the fosite interface.
|
||||
_, ar, err := provider.IntrospectToken(ctx, token, fosite.AccessToken, &oauthIntrospectSession{})
|
||||
// Decouple from the HTTP request context so token introspection completes
|
||||
// even if the client disconnects (fixes "context canceled" errors during
|
||||
// concurrent MCP connections from claude.ai).
|
||||
dbCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
_, ar, err := provider.IntrospectToken(dbCtx, token, fosite.AccessToken, &oauthIntrospectSession{})
|
||||
if err != nil {
|
||||
return "", "", false
|
||||
}
|
||||
|
||||
@@ -245,6 +245,12 @@ func (s *AgentService) ListAgents(ctx context.Context, ownerID int64) ([]*Agent,
|
||||
return s.store.ListAgentsByOwner(ctx, ownerID)
|
||||
}
|
||||
|
||||
// ListAllActiveAgents returns all active non-human agents across all owners.
|
||||
// Used for the A2A Agent Card discovery endpoint.
|
||||
func (s *AgentService) ListAllActiveAgents(ctx context.Context) ([]*Agent, error) {
|
||||
return s.store.ListAllActiveAgents(ctx)
|
||||
}
|
||||
|
||||
// RevokeKey generates a new API key for an agent. Only the owner can do this.
|
||||
// Returns the agent and the new raw API key (shown once).
|
||||
func (s *AgentService) RevokeKey(ctx context.Context, name string, ownerID int64) (*Agent, string, error) {
|
||||
|
||||
@@ -15,6 +15,7 @@ type AgentStore interface {
|
||||
UpdateAgent(ctx context.Context, agent *Agent) error
|
||||
DeactivateAgent(ctx context.Context, name string) error
|
||||
ListActiveAgents(ctx context.Context) ([]*Agent, error)
|
||||
ListAllActiveAgents(ctx context.Context) ([]*Agent, error)
|
||||
ListAgentsByOwner(ctx context.Context, ownerID int64) ([]*Agent, error)
|
||||
SearchAgentsByCapability(ctx context.Context, query string) ([]*Agent, error)
|
||||
GetHumanAgentByOwner(ctx context.Context, ownerID int64) (*Agent, error)
|
||||
@@ -112,6 +113,18 @@ func (s *SQLiteAgentStore) ListActiveAgents(ctx context.Context) ([]*Agent, erro
|
||||
return s.scanAgents(rows)
|
||||
}
|
||||
|
||||
func (s *SQLiteAgentStore) ListAllActiveAgents(ctx context.Context) ([]*Agent, error) {
|
||||
rows, err := s.db.QueryContext(ctx,
|
||||
`SELECT id, name, display_name, type, capabilities, owner_id, api_key_hash, status, created_at, updated_at
|
||||
FROM agents WHERE status = 'active' AND type != 'human' ORDER BY name`,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
return s.scanAgents(rows)
|
||||
}
|
||||
|
||||
func (s *SQLiteAgentStore) ListAgentsByOwner(ctx context.Context, ownerID int64) ([]*Agent, error) {
|
||||
rows, err := s.db.QueryContext(ctx,
|
||||
`SELECT id, name, display_name, type, capabilities, owner_id, api_key_hash, status, created_at, updated_at
|
||||
|
||||
@@ -0,0 +1,282 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/synapbus/synapbus/internal/agents"
|
||||
"github.com/synapbus/synapbus/internal/channels"
|
||||
)
|
||||
|
||||
// AnalyticsHandler serves analytics endpoints for the Web UI dashboard.
|
||||
type AnalyticsHandler struct {
|
||||
db *sql.DB
|
||||
agentService *agents.AgentService
|
||||
channelService *channels.Service
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
// NewAnalyticsHandler creates a new analytics handler.
|
||||
func NewAnalyticsHandler(db *sql.DB, agentService *agents.AgentService, channelService *channels.Service) *AnalyticsHandler {
|
||||
return &AnalyticsHandler{
|
||||
db: db,
|
||||
agentService: agentService,
|
||||
channelService: channelService,
|
||||
logger: slog.Default().With("component", "api.analytics"),
|
||||
}
|
||||
}
|
||||
|
||||
// parseSpan parses a span parameter and returns the cutoff time and strftime format string.
|
||||
func parseSpan(span string) (time.Time, string, error) {
|
||||
now := time.Now().UTC()
|
||||
switch span {
|
||||
case "1h":
|
||||
return now.Add(-1 * time.Hour), "%Y-%m-%d %H:%M", nil
|
||||
case "4h":
|
||||
return now.Add(-4 * time.Hour), "%Y-%m-%d %H:%M", nil
|
||||
case "24h":
|
||||
return now.Add(-24 * time.Hour), "%Y-%m-%d %H:00", nil
|
||||
case "7d":
|
||||
return now.Add(-7 * 24 * time.Hour), "%Y-%m-%d", nil
|
||||
case "30d":
|
||||
return now.Add(-30 * 24 * time.Hour), "%Y-%m-%d", nil
|
||||
default:
|
||||
return time.Time{}, "", fmt.Errorf("invalid span: %s (valid: 1h, 4h, 24h, 7d, 30d)", span)
|
||||
}
|
||||
}
|
||||
|
||||
// Timeline handles GET /api/analytics/timeline?span=24h.
|
||||
func (h *AnalyticsHandler) Timeline(w http.ResponseWriter, r *http.Request) {
|
||||
span := r.URL.Query().Get("span")
|
||||
if span == "" {
|
||||
span = "24h"
|
||||
}
|
||||
|
||||
cutoff, strftimeFmt, err := parseSpan(span)
|
||||
if err != nil {
|
||||
writeJSON(w, http.StatusBadRequest, errorBody("invalid_span", err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
// For 1h and 4h spans, bucket by 5-minute intervals using SQL expression
|
||||
var query string
|
||||
if span == "1h" || span == "4h" {
|
||||
// Round minutes down to nearest 5 using printf
|
||||
query = `SELECT strftime('%Y-%m-%d %H:', created_at) || printf('%02d', (CAST(strftime('%M', created_at) AS INTEGER) / 5) * 5) AS bucket, COUNT(*) AS count FROM messages WHERE created_at >= ? GROUP BY bucket ORDER BY bucket`
|
||||
} else {
|
||||
query = `SELECT strftime(?, created_at) AS bucket, COUNT(*) AS count FROM messages WHERE created_at >= ? GROUP BY bucket ORDER BY bucket`
|
||||
}
|
||||
|
||||
type bucket struct {
|
||||
Time string `json:"time"`
|
||||
Count int `json:"count"`
|
||||
}
|
||||
|
||||
var buckets []bucket
|
||||
var total int
|
||||
|
||||
cutoffStr := cutoff.Format("2006-01-02 15:04:05")
|
||||
|
||||
var rows *sql.Rows
|
||||
if span == "1h" || span == "4h" {
|
||||
rows, err = h.db.QueryContext(r.Context(), query, cutoffStr)
|
||||
} else {
|
||||
rows, err = h.db.QueryContext(r.Context(), query, strftimeFmt, cutoffStr)
|
||||
}
|
||||
if err != nil {
|
||||
h.logger.Error("timeline query failed", "error", err)
|
||||
writeJSON(w, http.StatusInternalServerError, errorBody("server_error", "Failed to query timeline"))
|
||||
return
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
for rows.Next() {
|
||||
var b bucket
|
||||
if err := rows.Scan(&b.Time, &b.Count); err != nil {
|
||||
h.logger.Error("timeline scan failed", "error", err)
|
||||
continue
|
||||
}
|
||||
buckets = append(buckets, b)
|
||||
total += b.Count
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
h.logger.Error("timeline rows error", "error", err)
|
||||
}
|
||||
|
||||
if buckets == nil {
|
||||
buckets = []bucket{}
|
||||
}
|
||||
|
||||
writeJSON(w, http.StatusOK, map[string]any{
|
||||
"span": span,
|
||||
"buckets": buckets,
|
||||
"total": total,
|
||||
})
|
||||
}
|
||||
|
||||
// TopAgents handles GET /api/analytics/top-agents?span=24h&limit=5.
|
||||
func (h *AnalyticsHandler) TopAgents(w http.ResponseWriter, r *http.Request) {
|
||||
span := r.URL.Query().Get("span")
|
||||
if span == "" {
|
||||
span = "24h"
|
||||
}
|
||||
|
||||
cutoff, _, err := parseSpan(span)
|
||||
if err != nil {
|
||||
writeJSON(w, http.StatusBadRequest, errorBody("invalid_span", err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
limit := 5
|
||||
if l := r.URL.Query().Get("limit"); l != "" {
|
||||
if parsed, err := strconv.Atoi(l); err == nil && parsed > 0 {
|
||||
limit = parsed
|
||||
if limit > 100 {
|
||||
limit = 100
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
cutoffStr := cutoff.Format("2006-01-02 15:04:05")
|
||||
|
||||
rows, err := h.db.QueryContext(r.Context(),
|
||||
`SELECT from_agent, COUNT(*) as count FROM messages WHERE created_at >= ? AND from_agent != '' GROUP BY from_agent ORDER BY count DESC LIMIT ?`,
|
||||
cutoffStr, limit,
|
||||
)
|
||||
if err != nil {
|
||||
h.logger.Error("top-agents query failed", "error", err)
|
||||
writeJSON(w, http.StatusInternalServerError, errorBody("server_error", "Failed to query top agents"))
|
||||
return
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
type agentStat struct {
|
||||
Name string `json:"name"`
|
||||
DisplayName string `json:"display_name"`
|
||||
Count int `json:"count"`
|
||||
}
|
||||
|
||||
var agentStats []agentStat
|
||||
for rows.Next() {
|
||||
var s agentStat
|
||||
if err := rows.Scan(&s.Name, &s.Count); err != nil {
|
||||
h.logger.Error("top-agents scan failed", "error", err)
|
||||
continue
|
||||
}
|
||||
|
||||
// Look up display name from agent service
|
||||
if h.agentService != nil {
|
||||
if agent, err := h.agentService.GetAgent(r.Context(), s.Name); err == nil {
|
||||
s.DisplayName = agent.DisplayName
|
||||
}
|
||||
}
|
||||
if s.DisplayName == "" {
|
||||
s.DisplayName = s.Name
|
||||
}
|
||||
|
||||
agentStats = append(agentStats, s)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
h.logger.Error("top-agents rows error", "error", err)
|
||||
}
|
||||
|
||||
if agentStats == nil {
|
||||
agentStats = []agentStat{}
|
||||
}
|
||||
|
||||
writeJSON(w, http.StatusOK, map[string]any{
|
||||
"span": span,
|
||||
"agents": agentStats,
|
||||
})
|
||||
}
|
||||
|
||||
// TopChannels handles GET /api/analytics/top-channels?span=24h&limit=5.
|
||||
func (h *AnalyticsHandler) TopChannels(w http.ResponseWriter, r *http.Request) {
|
||||
span := r.URL.Query().Get("span")
|
||||
if span == "" {
|
||||
span = "24h"
|
||||
}
|
||||
|
||||
cutoff, _, err := parseSpan(span)
|
||||
if err != nil {
|
||||
writeJSON(w, http.StatusBadRequest, errorBody("invalid_span", err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
limit := 5
|
||||
if l := r.URL.Query().Get("limit"); l != "" {
|
||||
if parsed, err := strconv.Atoi(l); err == nil && parsed > 0 {
|
||||
limit = parsed
|
||||
if limit > 100 {
|
||||
limit = 100
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
cutoffStr := cutoff.Format("2006-01-02 15:04:05")
|
||||
|
||||
rows, err := h.db.QueryContext(r.Context(),
|
||||
`SELECT c.name, COUNT(*) as count FROM messages m JOIN channels c ON m.channel_id = c.id WHERE m.created_at >= ? AND m.channel_id IS NOT NULL GROUP BY c.name ORDER BY count DESC LIMIT ?`,
|
||||
cutoffStr, limit,
|
||||
)
|
||||
if err != nil {
|
||||
h.logger.Error("top-channels query failed", "error", err)
|
||||
writeJSON(w, http.StatusInternalServerError, errorBody("server_error", "Failed to query top channels"))
|
||||
return
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
type channelStat struct {
|
||||
Name string `json:"name"`
|
||||
Count int `json:"count"`
|
||||
}
|
||||
|
||||
var channelStats []channelStat
|
||||
for rows.Next() {
|
||||
var s channelStat
|
||||
if err := rows.Scan(&s.Name, &s.Count); err != nil {
|
||||
h.logger.Error("top-channels scan failed", "error", err)
|
||||
continue
|
||||
}
|
||||
channelStats = append(channelStats, s)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
h.logger.Error("top-channels rows error", "error", err)
|
||||
}
|
||||
|
||||
if channelStats == nil {
|
||||
channelStats = []channelStat{}
|
||||
}
|
||||
|
||||
writeJSON(w, http.StatusOK, map[string]any{
|
||||
"span": span,
|
||||
"channels": channelStats,
|
||||
})
|
||||
}
|
||||
|
||||
// Summary handles GET /api/analytics/summary.
|
||||
func (h *AnalyticsHandler) Summary(w http.ResponseWriter, r *http.Request) {
|
||||
var totalAgents, totalChannels, totalMessages int
|
||||
|
||||
if err := h.db.QueryRowContext(r.Context(), `SELECT COUNT(*) FROM agents WHERE status = 'active'`).Scan(&totalAgents); err != nil {
|
||||
h.logger.Error("summary agents count failed", "error", err)
|
||||
}
|
||||
|
||||
if err := h.db.QueryRowContext(r.Context(), `SELECT COUNT(*) FROM channels`).Scan(&totalChannels); err != nil {
|
||||
h.logger.Error("summary channels count failed", "error", err)
|
||||
}
|
||||
|
||||
if err := h.db.QueryRowContext(r.Context(), `SELECT COUNT(*) FROM messages`).Scan(&totalMessages); err != nil {
|
||||
h.logger.Error("summary messages count failed", "error", err)
|
||||
}
|
||||
|
||||
writeJSON(w, http.StatusOK, map[string]any{
|
||||
"total_agents": totalAgents,
|
||||
"total_channels": totalChannels,
|
||||
"total_messages": totalMessages,
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,556 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/go-chi/chi/v5"
|
||||
_ "modernc.org/sqlite"
|
||||
|
||||
"github.com/synapbus/synapbus/internal/agents"
|
||||
"github.com/synapbus/synapbus/internal/channels"
|
||||
"github.com/synapbus/synapbus/internal/messaging"
|
||||
"github.com/synapbus/synapbus/internal/storage"
|
||||
)
|
||||
|
||||
func newAnalyticsTestDB(t *testing.T) *sql.DB {
|
||||
t.Helper()
|
||||
dsn := fmt.Sprintf("file:analytics_%s?mode=memory&cache=shared", t.Name())
|
||||
db, err := sql.Open("sqlite", dsn)
|
||||
if err != nil {
|
||||
t.Fatalf("open database: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { db.Close() })
|
||||
|
||||
if _, err := db.Exec("PRAGMA foreign_keys=ON"); err != nil {
|
||||
t.Fatalf("enable foreign keys: %v", err)
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
if err := storage.RunMigrations(ctx, db); err != nil {
|
||||
t.Fatalf("run migrations: %v", err)
|
||||
}
|
||||
|
||||
// Seed test users
|
||||
db.Exec(`INSERT OR IGNORE INTO users (id, username, password_hash, display_name) VALUES (1, 'testuser', 'hash', 'Test User')`)
|
||||
|
||||
return db
|
||||
}
|
||||
|
||||
func seedAnalyticsAgent(t *testing.T, db *sql.DB, name, displayName string, ownerID int64) {
|
||||
t.Helper()
|
||||
_, err := db.Exec(
|
||||
`INSERT OR IGNORE INTO agents (name, display_name, type, capabilities, owner_id, api_key_hash, status) VALUES (?, ?, 'ai', '{}', ?, 'testhash', 'active')`,
|
||||
name, displayName, ownerID,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("seed agent %s: %v", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
func seedAnalyticsChannel(t *testing.T, db *sql.DB, name, createdBy string) int64 {
|
||||
t.Helper()
|
||||
result, err := db.Exec(
|
||||
`INSERT INTO channels (name, description, type, is_private, created_by, created_at, updated_at) VALUES (?, '', 'standard', 0, ?, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`,
|
||||
name, createdBy,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("seed channel %s: %v", name, err)
|
||||
}
|
||||
id, _ := result.LastInsertId()
|
||||
return id
|
||||
}
|
||||
|
||||
func seedAnalyticsMessage(t *testing.T, db *sql.DB, from, to, body string, channelID *int64, createdAt time.Time) {
|
||||
t.Helper()
|
||||
|
||||
// Ensure a conversation exists
|
||||
result, err := db.Exec(
|
||||
`INSERT INTO conversations (subject, created_by, created_at, updated_at) VALUES ('test', ?, ?, ?)`,
|
||||
from, createdAt.Format("2006-01-02 15:04:05"), createdAt.Format("2006-01-02 15:04:05"),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("seed conversation: %v", err)
|
||||
}
|
||||
convID, _ := result.LastInsertId()
|
||||
|
||||
_, err = db.Exec(
|
||||
`INSERT INTO messages (conversation_id, from_agent, to_agent, channel_id, body, priority, status, metadata, created_at, updated_at) VALUES (?, ?, ?, ?, ?, 5, 'pending', '{}', ?, ?)`,
|
||||
convID, from, to, channelID, body, createdAt.Format("2006-01-02 15:04:05"), createdAt.Format("2006-01-02 15:04:05"),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("seed message: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func setupAnalyticsRouter(t *testing.T, db *sql.DB) chi.Router {
|
||||
t.Helper()
|
||||
|
||||
agentStore := agents.NewSQLiteAgentStore(db)
|
||||
agentService := agents.NewAgentService(agentStore, nil)
|
||||
|
||||
channelStore := channels.NewSQLiteChannelStore(db)
|
||||
msgStore := messaging.NewSQLiteMessageStore(db)
|
||||
msgService := messaging.NewMessagingService(msgStore, nil)
|
||||
channelService := channels.NewService(channelStore, msgService, nil)
|
||||
|
||||
analyticsHandler := NewAnalyticsHandler(db, agentService, channelService)
|
||||
|
||||
router := chi.NewRouter()
|
||||
router.Group(func(r chi.Router) {
|
||||
r.Use(OwnerAuthMiddleware)
|
||||
r.Get("/api/analytics/timeline", analyticsHandler.Timeline)
|
||||
r.Get("/api/analytics/top-agents", analyticsHandler.TopAgents)
|
||||
r.Get("/api/analytics/top-channels", analyticsHandler.TopChannels)
|
||||
r.Get("/api/analytics/summary", analyticsHandler.Summary)
|
||||
})
|
||||
|
||||
return router
|
||||
}
|
||||
|
||||
func analyticsRequest(t *testing.T, router chi.Router, method, path string) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
req := httptest.NewRequest(method, path, nil)
|
||||
req.Header.Set("X-Owner-ID", "1")
|
||||
rr := httptest.NewRecorder()
|
||||
router.ServeHTTP(rr, req)
|
||||
return rr
|
||||
}
|
||||
|
||||
func TestAnalyticsTimeline(t *testing.T) {
|
||||
db := newAnalyticsTestDB(t)
|
||||
seedAnalyticsAgent(t, db, "agent-a", "Agent A", 1)
|
||||
seedAnalyticsAgent(t, db, "agent-b", "Agent B", 1)
|
||||
|
||||
now := time.Now().UTC()
|
||||
|
||||
// Seed messages at various times within the last 24h
|
||||
for i := 0; i < 5; i++ {
|
||||
seedAnalyticsMessage(t, db, "agent-a", "agent-b", fmt.Sprintf("msg-%d", i), nil, now.Add(-time.Duration(i)*time.Hour))
|
||||
}
|
||||
|
||||
router := setupAnalyticsRouter(t, db)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
span string
|
||||
wantStatus int
|
||||
wantNonEmpty bool
|
||||
}{
|
||||
{
|
||||
name: "default span (24h)",
|
||||
span: "",
|
||||
wantStatus: http.StatusOK,
|
||||
wantNonEmpty: true,
|
||||
},
|
||||
{
|
||||
name: "1h span",
|
||||
span: "1h",
|
||||
wantStatus: http.StatusOK,
|
||||
wantNonEmpty: true,
|
||||
},
|
||||
{
|
||||
name: "4h span",
|
||||
span: "4h",
|
||||
wantStatus: http.StatusOK,
|
||||
wantNonEmpty: true,
|
||||
},
|
||||
{
|
||||
name: "24h span",
|
||||
span: "24h",
|
||||
wantStatus: http.StatusOK,
|
||||
wantNonEmpty: true,
|
||||
},
|
||||
{
|
||||
name: "7d span",
|
||||
span: "7d",
|
||||
wantStatus: http.StatusOK,
|
||||
wantNonEmpty: true,
|
||||
},
|
||||
{
|
||||
name: "30d span",
|
||||
span: "30d",
|
||||
wantStatus: http.StatusOK,
|
||||
wantNonEmpty: true,
|
||||
},
|
||||
{
|
||||
name: "invalid span",
|
||||
span: "99x",
|
||||
wantStatus: http.StatusBadRequest,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
path := "/api/analytics/timeline"
|
||||
if tt.span != "" {
|
||||
path += "?span=" + tt.span
|
||||
}
|
||||
|
||||
rr := analyticsRequest(t, router, "GET", path)
|
||||
if rr.Code != tt.wantStatus {
|
||||
t.Fatalf("status = %d, want %d, body: %s", rr.Code, tt.wantStatus, rr.Body.String())
|
||||
}
|
||||
|
||||
if tt.wantStatus != http.StatusOK {
|
||||
return
|
||||
}
|
||||
|
||||
var resp struct {
|
||||
Span string `json:"span"`
|
||||
Buckets []struct {
|
||||
Time string `json:"time"`
|
||||
Count int `json:"count"`
|
||||
} `json:"buckets"`
|
||||
Total int `json:"total"`
|
||||
}
|
||||
if err := json.Unmarshal(rr.Body.Bytes(), &resp); err != nil {
|
||||
t.Fatalf("unmarshal: %v", err)
|
||||
}
|
||||
|
||||
expectedSpan := tt.span
|
||||
if expectedSpan == "" {
|
||||
expectedSpan = "24h"
|
||||
}
|
||||
if resp.Span != expectedSpan {
|
||||
t.Errorf("span = %q, want %q", resp.Span, expectedSpan)
|
||||
}
|
||||
|
||||
if tt.wantNonEmpty && resp.Total == 0 {
|
||||
t.Error("expected non-zero total")
|
||||
}
|
||||
|
||||
// Verify total matches sum of bucket counts
|
||||
bucketSum := 0
|
||||
for _, b := range resp.Buckets {
|
||||
bucketSum += b.Count
|
||||
}
|
||||
if bucketSum != resp.Total {
|
||||
t.Errorf("bucket sum %d != total %d", bucketSum, resp.Total)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestAnalyticsTimeline_EmptyDB(t *testing.T) {
|
||||
db := newAnalyticsTestDB(t)
|
||||
router := setupAnalyticsRouter(t, db)
|
||||
|
||||
rr := analyticsRequest(t, router, "GET", "/api/analytics/timeline?span=24h")
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want %d", rr.Code, http.StatusOK)
|
||||
}
|
||||
|
||||
var resp struct {
|
||||
Span string `json:"span"`
|
||||
Buckets []any `json:"buckets"`
|
||||
Total int `json:"total"`
|
||||
}
|
||||
json.Unmarshal(rr.Body.Bytes(), &resp)
|
||||
|
||||
if resp.Total != 0 {
|
||||
t.Errorf("total = %d, want 0", resp.Total)
|
||||
}
|
||||
if len(resp.Buckets) != 0 {
|
||||
t.Errorf("buckets = %d, want 0", len(resp.Buckets))
|
||||
}
|
||||
}
|
||||
|
||||
func TestAnalyticsTopAgents(t *testing.T) {
|
||||
db := newAnalyticsTestDB(t)
|
||||
seedAnalyticsAgent(t, db, "agent-alpha", "Alpha Agent", 1)
|
||||
seedAnalyticsAgent(t, db, "agent-beta", "Beta Agent", 1)
|
||||
|
||||
now := time.Now().UTC()
|
||||
|
||||
// agent-alpha sends 5 messages, agent-beta sends 2
|
||||
for i := 0; i < 5; i++ {
|
||||
seedAnalyticsMessage(t, db, "agent-alpha", "agent-beta", fmt.Sprintf("msg-%d", i), nil, now.Add(-time.Duration(i)*time.Minute))
|
||||
}
|
||||
for i := 0; i < 2; i++ {
|
||||
seedAnalyticsMessage(t, db, "agent-beta", "agent-alpha", fmt.Sprintf("reply-%d", i), nil, now.Add(-time.Duration(i)*time.Minute))
|
||||
}
|
||||
|
||||
router := setupAnalyticsRouter(t, db)
|
||||
|
||||
t.Run("returns agents sorted by count", func(t *testing.T) {
|
||||
rr := analyticsRequest(t, router, "GET", "/api/analytics/top-agents?span=24h")
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want %d, body: %s", rr.Code, http.StatusOK, rr.Body.String())
|
||||
}
|
||||
|
||||
var resp struct {
|
||||
Span string `json:"span"`
|
||||
Agents []struct {
|
||||
Name string `json:"name"`
|
||||
DisplayName string `json:"display_name"`
|
||||
Count int `json:"count"`
|
||||
} `json:"agents"`
|
||||
}
|
||||
json.Unmarshal(rr.Body.Bytes(), &resp)
|
||||
|
||||
if len(resp.Agents) < 2 {
|
||||
t.Fatalf("agents = %d, want >= 2", len(resp.Agents))
|
||||
}
|
||||
|
||||
// First agent should be alpha (5 messages)
|
||||
if resp.Agents[0].Name != "agent-alpha" {
|
||||
t.Errorf("top agent = %q, want agent-alpha", resp.Agents[0].Name)
|
||||
}
|
||||
if resp.Agents[0].Count != 5 {
|
||||
t.Errorf("top agent count = %d, want 5", resp.Agents[0].Count)
|
||||
}
|
||||
if resp.Agents[0].DisplayName != "Alpha Agent" {
|
||||
t.Errorf("display_name = %q, want 'Alpha Agent'", resp.Agents[0].DisplayName)
|
||||
}
|
||||
|
||||
// Second agent should be beta (2 messages)
|
||||
if resp.Agents[1].Name != "agent-beta" {
|
||||
t.Errorf("second agent = %q, want agent-beta", resp.Agents[1].Name)
|
||||
}
|
||||
if resp.Agents[1].Count != 2 {
|
||||
t.Errorf("second agent count = %d, want 2", resp.Agents[1].Count)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("respects limit parameter", func(t *testing.T) {
|
||||
rr := analyticsRequest(t, router, "GET", "/api/analytics/top-agents?span=24h&limit=1")
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want %d", rr.Code, http.StatusOK)
|
||||
}
|
||||
|
||||
var resp struct {
|
||||
Agents []struct {
|
||||
Name string `json:"name"`
|
||||
Count int `json:"count"`
|
||||
} `json:"agents"`
|
||||
}
|
||||
json.Unmarshal(rr.Body.Bytes(), &resp)
|
||||
|
||||
if len(resp.Agents) != 1 {
|
||||
t.Errorf("agents = %d, want 1", len(resp.Agents))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("invalid span returns error", func(t *testing.T) {
|
||||
rr := analyticsRequest(t, router, "GET", "/api/analytics/top-agents?span=invalid")
|
||||
if rr.Code != http.StatusBadRequest {
|
||||
t.Errorf("status = %d, want %d", rr.Code, http.StatusBadRequest)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestAnalyticsTopAgents_EmptyDB(t *testing.T) {
|
||||
db := newAnalyticsTestDB(t)
|
||||
router := setupAnalyticsRouter(t, db)
|
||||
|
||||
rr := analyticsRequest(t, router, "GET", "/api/analytics/top-agents?span=24h")
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want %d", rr.Code, http.StatusOK)
|
||||
}
|
||||
|
||||
var resp struct {
|
||||
Agents []any `json:"agents"`
|
||||
}
|
||||
json.Unmarshal(rr.Body.Bytes(), &resp)
|
||||
|
||||
if len(resp.Agents) != 0 {
|
||||
t.Errorf("agents = %d, want 0", len(resp.Agents))
|
||||
}
|
||||
}
|
||||
|
||||
func TestAnalyticsTopChannels(t *testing.T) {
|
||||
db := newAnalyticsTestDB(t)
|
||||
seedAnalyticsAgent(t, db, "agent-a", "Agent A", 1)
|
||||
|
||||
chID1 := seedAnalyticsChannel(t, db, "news-mcp", "agent-a")
|
||||
chID2 := seedAnalyticsChannel(t, db, "general", "agent-a")
|
||||
|
||||
now := time.Now().UTC()
|
||||
|
||||
// news-mcp gets 4 messages, general gets 2
|
||||
for i := 0; i < 4; i++ {
|
||||
seedAnalyticsMessage(t, db, "agent-a", "", fmt.Sprintf("ch-msg-%d", i), &chID1, now.Add(-time.Duration(i)*time.Minute))
|
||||
}
|
||||
for i := 0; i < 2; i++ {
|
||||
seedAnalyticsMessage(t, db, "agent-a", "", fmt.Sprintf("gen-msg-%d", i), &chID2, now.Add(-time.Duration(i)*time.Minute))
|
||||
}
|
||||
|
||||
router := setupAnalyticsRouter(t, db)
|
||||
|
||||
t.Run("returns channels sorted by count", func(t *testing.T) {
|
||||
rr := analyticsRequest(t, router, "GET", "/api/analytics/top-channels?span=24h")
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want %d, body: %s", rr.Code, http.StatusOK, rr.Body.String())
|
||||
}
|
||||
|
||||
var resp struct {
|
||||
Span string `json:"span"`
|
||||
Channels []struct {
|
||||
Name string `json:"name"`
|
||||
Count int `json:"count"`
|
||||
} `json:"channels"`
|
||||
}
|
||||
json.Unmarshal(rr.Body.Bytes(), &resp)
|
||||
|
||||
if len(resp.Channels) < 2 {
|
||||
t.Fatalf("channels = %d, want >= 2", len(resp.Channels))
|
||||
}
|
||||
|
||||
if resp.Channels[0].Name != "news-mcp" {
|
||||
t.Errorf("top channel = %q, want news-mcp", resp.Channels[0].Name)
|
||||
}
|
||||
if resp.Channels[0].Count != 4 {
|
||||
t.Errorf("top channel count = %d, want 4", resp.Channels[0].Count)
|
||||
}
|
||||
|
||||
if resp.Channels[1].Name != "general" {
|
||||
t.Errorf("second channel = %q, want general", resp.Channels[1].Name)
|
||||
}
|
||||
if resp.Channels[1].Count != 2 {
|
||||
t.Errorf("second channel count = %d, want 2", resp.Channels[1].Count)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("respects limit parameter", func(t *testing.T) {
|
||||
rr := analyticsRequest(t, router, "GET", "/api/analytics/top-channels?span=24h&limit=1")
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want %d", rr.Code, http.StatusOK)
|
||||
}
|
||||
|
||||
var resp struct {
|
||||
Channels []struct {
|
||||
Name string `json:"name"`
|
||||
Count int `json:"count"`
|
||||
} `json:"channels"`
|
||||
}
|
||||
json.Unmarshal(rr.Body.Bytes(), &resp)
|
||||
|
||||
if len(resp.Channels) != 1 {
|
||||
t.Errorf("channels = %d, want 1", len(resp.Channels))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("invalid span returns error", func(t *testing.T) {
|
||||
rr := analyticsRequest(t, router, "GET", "/api/analytics/top-channels?span=invalid")
|
||||
if rr.Code != http.StatusBadRequest {
|
||||
t.Errorf("status = %d, want %d", rr.Code, http.StatusBadRequest)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestAnalyticsTopChannels_EmptyDB(t *testing.T) {
|
||||
db := newAnalyticsTestDB(t)
|
||||
router := setupAnalyticsRouter(t, db)
|
||||
|
||||
rr := analyticsRequest(t, router, "GET", "/api/analytics/top-channels?span=24h")
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want %d", rr.Code, http.StatusOK)
|
||||
}
|
||||
|
||||
var resp struct {
|
||||
Channels []any `json:"channels"`
|
||||
}
|
||||
json.Unmarshal(rr.Body.Bytes(), &resp)
|
||||
|
||||
if len(resp.Channels) != 0 {
|
||||
t.Errorf("channels = %d, want 0", len(resp.Channels))
|
||||
}
|
||||
}
|
||||
|
||||
func TestAnalyticsSummary(t *testing.T) {
|
||||
db := newAnalyticsTestDB(t)
|
||||
seedAnalyticsAgent(t, db, "agent-a", "Agent A", 1)
|
||||
seedAnalyticsAgent(t, db, "agent-b", "Agent B", 1)
|
||||
|
||||
chID := seedAnalyticsChannel(t, db, "test-channel", "agent-a")
|
||||
|
||||
now := time.Now().UTC()
|
||||
seedAnalyticsMessage(t, db, "agent-a", "agent-b", "hello", nil, now)
|
||||
seedAnalyticsMessage(t, db, "agent-b", "agent-a", "hi back", nil, now)
|
||||
seedAnalyticsMessage(t, db, "agent-a", "", "channel msg", &chID, now)
|
||||
|
||||
router := setupAnalyticsRouter(t, db)
|
||||
|
||||
rr := analyticsRequest(t, router, "GET", "/api/analytics/summary")
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want %d, body: %s", rr.Code, http.StatusOK, rr.Body.String())
|
||||
}
|
||||
|
||||
var resp struct {
|
||||
TotalAgents int `json:"total_agents"`
|
||||
TotalChannels int `json:"total_channels"`
|
||||
TotalMessages int `json:"total_messages"`
|
||||
}
|
||||
if err := json.Unmarshal(rr.Body.Bytes(), &resp); err != nil {
|
||||
t.Fatalf("unmarshal: %v", err)
|
||||
}
|
||||
|
||||
if resp.TotalAgents != 2 {
|
||||
t.Errorf("total_agents = %d, want 2", resp.TotalAgents)
|
||||
}
|
||||
if resp.TotalChannels != 1 {
|
||||
t.Errorf("total_channels = %d, want 1", resp.TotalChannels)
|
||||
}
|
||||
if resp.TotalMessages != 3 {
|
||||
t.Errorf("total_messages = %d, want 3", resp.TotalMessages)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAnalyticsSummary_EmptyDB(t *testing.T) {
|
||||
db := newAnalyticsTestDB(t)
|
||||
router := setupAnalyticsRouter(t, db)
|
||||
|
||||
rr := analyticsRequest(t, router, "GET", "/api/analytics/summary")
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want %d", rr.Code, http.StatusOK)
|
||||
}
|
||||
|
||||
var resp struct {
|
||||
TotalAgents int `json:"total_agents"`
|
||||
TotalChannels int `json:"total_channels"`
|
||||
TotalMessages int `json:"total_messages"`
|
||||
}
|
||||
json.Unmarshal(rr.Body.Bytes(), &resp)
|
||||
|
||||
if resp.TotalAgents != 0 {
|
||||
t.Errorf("total_agents = %d, want 0", resp.TotalAgents)
|
||||
}
|
||||
if resp.TotalChannels != 0 {
|
||||
t.Errorf("total_channels = %d, want 0", resp.TotalChannels)
|
||||
}
|
||||
if resp.TotalMessages != 0 {
|
||||
t.Errorf("total_messages = %d, want 0", resp.TotalMessages)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAnalytics_Unauthenticated(t *testing.T) {
|
||||
db := newAnalyticsTestDB(t)
|
||||
router := setupAnalyticsRouter(t, db)
|
||||
|
||||
endpoints := []string{
|
||||
"/api/analytics/timeline",
|
||||
"/api/analytics/top-agents",
|
||||
"/api/analytics/top-channels",
|
||||
"/api/analytics/summary",
|
||||
}
|
||||
|
||||
for _, endpoint := range endpoints {
|
||||
t.Run(endpoint, func(t *testing.T) {
|
||||
req := httptest.NewRequest("GET", endpoint, nil)
|
||||
// No X-Owner-ID header
|
||||
rr := httptest.NewRecorder()
|
||||
router.ServeHTTP(rr, req)
|
||||
|
||||
if rr.Code != http.StatusUnauthorized {
|
||||
t.Errorf("status = %d, want %d", rr.Code, http.StatusUnauthorized)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -137,6 +137,8 @@ func (h *AttachmentsHandler) Upload(w http.ResponseWriter, r *http.Request) {
|
||||
http.Error(w, `{"error":"empty file not allowed"}`, http.StatusBadRequest)
|
||||
case attachments.ErrFileTooLarge:
|
||||
http.Error(w, `{"error":"file exceeds maximum size of 50MB"}`, http.StatusRequestEntityTooLarge)
|
||||
case attachments.ErrUnsupportedType:
|
||||
http.Error(w, `{"error":"unsupported file type: only images, PDFs, and text files are allowed"}`, http.StatusBadRequest)
|
||||
default:
|
||||
h.logger.Error("upload attachment failed", "error", err)
|
||||
http.Error(w, `{"error":"internal server error"}`, http.StatusInternalServerError)
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
|
||||
"github.com/synapbus/synapbus/internal/agents"
|
||||
"github.com/synapbus/synapbus/internal/channels"
|
||||
"github.com/synapbus/synapbus/internal/messaging"
|
||||
)
|
||||
|
||||
// NewMessageEvent is broadcast when a new message is sent.
|
||||
@@ -107,3 +108,24 @@ func (b *SSEBroadcaster) BroadcastChannelMessage(ctx context.Context, channelID
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// OnMessageSent implements messaging.MessageListener so the SSEBroadcaster
|
||||
// can be wired directly into the messaging service. This ensures SSE events
|
||||
// fire for messages sent via MCP (agents) as well as the REST API.
|
||||
func (b *SSEBroadcaster) OnMessageSent(ctx context.Context, msg *messaging.Message) {
|
||||
event := NewMessageEvent{
|
||||
MessageID: msg.ID,
|
||||
FromAgent: msg.FromAgent,
|
||||
ToAgent: msg.ToAgent,
|
||||
}
|
||||
|
||||
if msg.ChannelID != nil {
|
||||
ch, err := b.channelService.GetChannel(ctx, *msg.ChannelID)
|
||||
if err == nil {
|
||||
event.Channel = ch.Name
|
||||
}
|
||||
b.BroadcastChannelMessage(ctx, *msg.ChannelID, event)
|
||||
} else {
|
||||
b.BroadcastDM(ctx, event)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
@@ -15,10 +16,16 @@ import (
|
||||
|
||||
// ChannelsHandler handles REST API requests for channels.
|
||||
type ChannelsHandler struct {
|
||||
channelService *channels.Service
|
||||
agentService *agents.AgentService
|
||||
msgService *messaging.MessagingService
|
||||
logger *slog.Logger
|
||||
channelService *channels.Service
|
||||
agentService *agents.AgentService
|
||||
msgService *messaging.MessagingService
|
||||
reactionService ChannelReactionService
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
// ChannelReactionService is the subset of reactions.Service needed by ChannelsHandler.
|
||||
type ChannelReactionService interface {
|
||||
ListByState(ctx context.Context, channelID int64, state string) ([]int64, error)
|
||||
}
|
||||
|
||||
// NewChannelsHandler creates a new channels handler.
|
||||
@@ -31,6 +38,11 @@ func NewChannelsHandler(channelService *channels.Service, agentService *agents.A
|
||||
}
|
||||
}
|
||||
|
||||
// SetReactionService sets the reaction service for workflow state queries.
|
||||
func (h *ChannelsHandler) SetReactionService(svc ChannelReactionService) {
|
||||
h.reactionService = svc
|
||||
}
|
||||
|
||||
// ListChannels handles GET /api/channels.
|
||||
func (h *ChannelsHandler) ListChannels(w http.ResponseWriter, r *http.Request) {
|
||||
ownerID, ok := OwnerIDFromContext(r.Context())
|
||||
@@ -228,6 +240,9 @@ func (h *ChannelsHandler) ChannelMessages(w http.ResponseWriter, r *http.Request
|
||||
return
|
||||
}
|
||||
|
||||
// Enrich messages with reply counts and attachments.
|
||||
h.msgService.EnrichMessages(r.Context(), paginated.Messages)
|
||||
|
||||
// Compute last_read_message_id across owned agents
|
||||
var lastReadMessageID int64
|
||||
ownedAgents, err := h.agentService.ListAgents(r.Context(), ownerID)
|
||||
@@ -295,3 +310,112 @@ func (h *ChannelsHandler) LeaveChannel(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
writeJSON(w, http.StatusOK, map[string]string{"status": "left"})
|
||||
}
|
||||
|
||||
// UpdateSettings handles PUT /api/channels/{name}/settings.
|
||||
func (h *ChannelsHandler) UpdateSettings(w http.ResponseWriter, r *http.Request) {
|
||||
_, ok := OwnerIDFromContext(r.Context())
|
||||
if !ok {
|
||||
writeJSON(w, http.StatusUnauthorized, errorBody("unauthorized", "Authentication required"))
|
||||
return
|
||||
}
|
||||
|
||||
name := chi.URLParam(r, "name")
|
||||
ch, err := h.channelService.GetChannelByName(r.Context(), name)
|
||||
if err != nil {
|
||||
writeJSON(w, http.StatusNotFound, errorBody("not_found", "Channel not found"))
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
AutoApprove *bool `json:"auto_approve"`
|
||||
StalemateRemindAfter *string `json:"stalemate_remind_after"`
|
||||
StalemateEscalateAfter *string `json:"stalemate_escalate_after"`
|
||||
}
|
||||
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
writeJSON(w, http.StatusBadRequest, errorBody("invalid_request", "Invalid JSON body"))
|
||||
return
|
||||
}
|
||||
|
||||
settings := channels.ChannelSettings{
|
||||
AutoApprove: ch.AutoApprove,
|
||||
StalemateRemindAfter: ch.StalemateRemindAfter,
|
||||
StalemateEscalateAfter: ch.StalemateEscalateAfter,
|
||||
}
|
||||
|
||||
if req.AutoApprove != nil {
|
||||
settings.AutoApprove = *req.AutoApprove
|
||||
}
|
||||
if req.StalemateRemindAfter != nil {
|
||||
settings.StalemateRemindAfter = *req.StalemateRemindAfter
|
||||
}
|
||||
if req.StalemateEscalateAfter != nil {
|
||||
settings.StalemateEscalateAfter = *req.StalemateEscalateAfter
|
||||
}
|
||||
|
||||
updated, err := h.channelService.UpdateChannelSettings(r.Context(), ch.ID, settings)
|
||||
if err != nil {
|
||||
h.logger.Error("update channel settings failed", "error", err)
|
||||
writeJSON(w, http.StatusInternalServerError, errorBody("server_error", "Failed to update channel settings"))
|
||||
return
|
||||
}
|
||||
|
||||
writeJSON(w, http.StatusOK, map[string]any{"channel": updated})
|
||||
}
|
||||
|
||||
// ListByState handles GET /api/channels/{name}/messages/by-state?state=X.
|
||||
func (h *ChannelsHandler) ListByState(w http.ResponseWriter, r *http.Request) {
|
||||
_, ok := OwnerIDFromContext(r.Context())
|
||||
if !ok {
|
||||
writeJSON(w, http.StatusUnauthorized, errorBody("unauthorized", "Authentication required"))
|
||||
return
|
||||
}
|
||||
|
||||
if h.reactionService == nil {
|
||||
writeJSON(w, http.StatusServiceUnavailable, errorBody("unavailable", "Reactions service not configured"))
|
||||
return
|
||||
}
|
||||
|
||||
name := chi.URLParam(r, "name")
|
||||
ch, err := h.channelService.GetChannelByName(r.Context(), name)
|
||||
if err != nil {
|
||||
writeJSON(w, http.StatusNotFound, errorBody("not_found", "Channel not found"))
|
||||
return
|
||||
}
|
||||
|
||||
state := r.URL.Query().Get("state")
|
||||
if state == "" {
|
||||
writeJSON(w, http.StatusBadRequest, errorBody("missing_state", "Query parameter 'state' is required"))
|
||||
return
|
||||
}
|
||||
|
||||
ids, err := h.reactionService.ListByState(r.Context(), ch.ID, state)
|
||||
if err != nil {
|
||||
h.logger.Error("list by state failed", "error", err)
|
||||
writeJSON(w, http.StatusBadRequest, errorBody("invalid_state", err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
// Load messages by IDs
|
||||
var messages []*messaging.Message
|
||||
for _, id := range ids {
|
||||
msg, err := h.msgService.GetMessageByID(r.Context(), id)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
messages = append(messages, msg)
|
||||
}
|
||||
|
||||
if messages == nil {
|
||||
messages = []*messaging.Message{}
|
||||
}
|
||||
|
||||
// Enrich messages with reactions, reply counts, attachments
|
||||
h.msgService.EnrichMessages(r.Context(), messages)
|
||||
|
||||
writeJSON(w, http.StatusOK, map[string]any{
|
||||
"messages": messages,
|
||||
"state": state,
|
||||
"total": len(messages),
|
||||
})
|
||||
}
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/go-chi/chi/v5"
|
||||
@@ -85,6 +86,8 @@ func (h *MessagesHandler) ListMessages(w http.ResponseWriter, r *http.Request) {
|
||||
allMessages = []*messaging.Message{}
|
||||
}
|
||||
|
||||
h.msgService.EnrichMessages(r.Context(), allMessages)
|
||||
|
||||
sortMessagesByTime(allMessages)
|
||||
if len(allMessages) > limit {
|
||||
allMessages = allMessages[:limit]
|
||||
@@ -121,6 +124,8 @@ func (h *MessagesHandler) GetMessage(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
h.msgService.EnrichMessages(r.Context(), []*messaging.Message{msg})
|
||||
|
||||
writeJSON(w, http.StatusOK, msg)
|
||||
}
|
||||
|
||||
@@ -229,6 +234,8 @@ func (h *MessagesHandler) GetConversation(w http.ResponseWriter, r *http.Request
|
||||
return
|
||||
}
|
||||
|
||||
h.msgService.EnrichMessages(r.Context(), messages)
|
||||
|
||||
writeJSON(w, http.StatusOK, map[string]any{
|
||||
"conversation": conv,
|
||||
"messages": messages,
|
||||
@@ -244,14 +251,15 @@ func (h *MessagesHandler) SendMessage(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
var req struct {
|
||||
From string `json:"from"`
|
||||
To string `json:"to"`
|
||||
Body string `json:"body"`
|
||||
Priority int `json:"priority"`
|
||||
ChannelID *int64 `json:"channel_id,omitempty"`
|
||||
ConversationID *int64 `json:"conversation_id,omitempty"`
|
||||
Subject string `json:"subject,omitempty"`
|
||||
ReplyTo *int64 `json:"reply_to,omitempty"`
|
||||
From string `json:"from"`
|
||||
To string `json:"to"`
|
||||
Body string `json:"body"`
|
||||
Priority int `json:"priority"`
|
||||
ChannelID *int64 `json:"channel_id,omitempty"`
|
||||
ConversationID *int64 `json:"conversation_id,omitempty"`
|
||||
Subject string `json:"subject,omitempty"`
|
||||
ReplyTo *int64 `json:"reply_to,omitempty"`
|
||||
Attachments []string `json:"attachments,omitempty"`
|
||||
}
|
||||
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
@@ -306,6 +314,7 @@ func (h *MessagesHandler) SendMessage(w http.ResponseWriter, r *http.Request) {
|
||||
ConversationID: req.ConversationID,
|
||||
Subject: req.Subject,
|
||||
ReplyTo: req.ReplyTo,
|
||||
Attachments: req.Attachments,
|
||||
}
|
||||
|
||||
msg, err := h.msgService.SendMessage(r.Context(), req.From, req.To, req.Body, opts)
|
||||
@@ -315,23 +324,10 @@ func (h *MessagesHandler) SendMessage(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
// Broadcast real-time event to connected SSE clients
|
||||
if h.broadcaster != nil {
|
||||
event := NewMessageEvent{
|
||||
MessageID: msg.ID,
|
||||
FromAgent: msg.FromAgent,
|
||||
ToAgent: msg.ToAgent,
|
||||
}
|
||||
if msg.ChannelID != nil && h.broadcaster.channelService != nil {
|
||||
ch, chErr := h.broadcaster.channelService.GetChannel(r.Context(), *msg.ChannelID)
|
||||
if chErr == nil {
|
||||
event.Channel = ch.Name
|
||||
}
|
||||
h.broadcaster.BroadcastChannelMessage(r.Context(), *msg.ChannelID, event)
|
||||
} else {
|
||||
h.broadcaster.BroadcastDM(r.Context(), event)
|
||||
}
|
||||
}
|
||||
// SSE broadcast is handled by the MessageListener on the messaging
|
||||
// service, so it fires for both REST and MCP message paths.
|
||||
|
||||
h.msgService.EnrichMessages(r.Context(), []*messaging.Message{msg})
|
||||
|
||||
writeJSON(w, http.StatusCreated, msg)
|
||||
}
|
||||
@@ -414,9 +410,60 @@ func (h *MessagesHandler) SearchMessages(w http.ResponseWriter, r *http.Request)
|
||||
limit = 20
|
||||
}
|
||||
|
||||
// Parse advanced filter parameters
|
||||
channelFilter := r.URL.Query().Get("channel")
|
||||
agentFilter := r.URL.Query().Get("agent")
|
||||
afterFilter := r.URL.Query().Get("after")
|
||||
beforeFilter := r.URL.Query().Get("before")
|
||||
|
||||
var allMessages []*messaging.Message
|
||||
for _, agent := range ownedAgents {
|
||||
opts := messaging.SearchOptions{Limit: limit}
|
||||
opts := messaging.SearchOptions{
|
||||
Limit: limit,
|
||||
After: afterFilter,
|
||||
Before: beforeFilter,
|
||||
}
|
||||
|
||||
// Apply channel filters (comma-separated, prefix with - to exclude)
|
||||
if channelFilter != "" {
|
||||
channels := strings.Split(channelFilter, ",")
|
||||
var includeChannels []string
|
||||
var excludeChannels []string
|
||||
for _, ch := range channels {
|
||||
ch = strings.TrimSpace(ch)
|
||||
if ch == "" {
|
||||
continue
|
||||
}
|
||||
if strings.HasPrefix(ch, "-") {
|
||||
excludeChannels = append(excludeChannels, strings.TrimPrefix(ch, "-"))
|
||||
} else {
|
||||
includeChannels = append(includeChannels, ch)
|
||||
}
|
||||
}
|
||||
opts.Channels = includeChannels
|
||||
opts.ExcludeChannels = excludeChannels
|
||||
}
|
||||
|
||||
// Apply agent filters (comma-separated, prefix with - to exclude)
|
||||
if agentFilter != "" {
|
||||
agentNames := strings.Split(agentFilter, ",")
|
||||
var includeAgents []string
|
||||
var excludeAgents []string
|
||||
for _, a := range agentNames {
|
||||
a = strings.TrimSpace(a)
|
||||
if a == "" {
|
||||
continue
|
||||
}
|
||||
if strings.HasPrefix(a, "-") {
|
||||
excludeAgents = append(excludeAgents, strings.TrimPrefix(a, "-"))
|
||||
} else {
|
||||
includeAgents = append(includeAgents, a)
|
||||
}
|
||||
}
|
||||
opts.Agents = includeAgents
|
||||
opts.ExcludeAgents = excludeAgents
|
||||
}
|
||||
|
||||
result, err := h.msgService.SearchMessages(r.Context(), agent.Name, query, opts)
|
||||
if err != nil {
|
||||
continue
|
||||
@@ -428,6 +475,8 @@ func (h *MessagesHandler) SearchMessages(w http.ResponseWriter, r *http.Request)
|
||||
allMessages = []*messaging.Message{}
|
||||
}
|
||||
|
||||
h.msgService.EnrichMessages(r.Context(), allMessages)
|
||||
|
||||
writeJSON(w, http.StatusOK, map[string]any{
|
||||
"messages": allMessages,
|
||||
"query": query,
|
||||
@@ -468,6 +517,8 @@ func (h *MessagesHandler) GetReplies(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
h.msgService.EnrichMessages(r.Context(), replies)
|
||||
|
||||
writeJSON(w, http.StatusOK, map[string]any{
|
||||
"replies": replies,
|
||||
"total": len(replies),
|
||||
@@ -513,6 +564,8 @@ func (h *MessagesHandler) DMMessages(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
h.msgService.EnrichMessages(r.Context(), msgs)
|
||||
|
||||
// Include last_read_message_id for the human agent's DM with the peer
|
||||
lastRead, _ := h.msgService.GetLastReadForDM(r.Context(), agentNames, peerAgent)
|
||||
|
||||
|
||||
@@ -333,7 +333,7 @@ func TestChannelMessages_IncludesLastRead(t *testing.T) {
|
||||
}
|
||||
|
||||
// Broadcast messages
|
||||
msgs, err := channelService.BroadcastMessage(ctx, ch.ID, "human-agent", "Hello channel", 5, "")
|
||||
msgs, err := channelService.BroadcastMessage(ctx, ch.ID, "human-agent", "Hello channel", 5, "", nil, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("broadcast: %v", err)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,112 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
|
||||
"github.com/synapbus/synapbus/internal/push"
|
||||
)
|
||||
|
||||
// PushHandler manages Web Push notification subscription endpoints.
|
||||
type PushHandler struct {
|
||||
pushService *push.Service
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
// NewPushHandler creates a new push notification handler.
|
||||
func NewPushHandler(pushService *push.Service) *PushHandler {
|
||||
return &PushHandler{
|
||||
pushService: pushService,
|
||||
logger: slog.Default().With("component", "api.push"),
|
||||
}
|
||||
}
|
||||
|
||||
// subscribeRequest is the JSON body for POST /api/push/subscribe.
|
||||
type subscribeRequest struct {
|
||||
Endpoint string `json:"endpoint"`
|
||||
KeyP256dh string `json:"key_p256dh"`
|
||||
KeyAuth string `json:"key_auth"`
|
||||
}
|
||||
|
||||
// unsubscribeRequest is the JSON body for DELETE /api/push/subscribe.
|
||||
type unsubscribeRequest struct {
|
||||
Endpoint string `json:"endpoint"`
|
||||
}
|
||||
|
||||
// Subscribe handles POST /api/push/subscribe.
|
||||
// Registers a Web Push subscription for the authenticated user.
|
||||
func (h *PushHandler) Subscribe(w http.ResponseWriter, r *http.Request) {
|
||||
ownerID, ok := OwnerIDFromContext(r.Context())
|
||||
if !ok {
|
||||
writeJSON(w, http.StatusUnauthorized, errorBody("unauthorized", "Authentication required"))
|
||||
return
|
||||
}
|
||||
|
||||
var req subscribeRequest
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
writeJSON(w, http.StatusBadRequest, errorBody("invalid_request", "Invalid JSON body"))
|
||||
return
|
||||
}
|
||||
|
||||
if req.Endpoint == "" || req.KeyP256dh == "" || req.KeyAuth == "" {
|
||||
writeJSON(w, http.StatusBadRequest, errorBody("validation_error", "endpoint, key_p256dh, and key_auth are required"))
|
||||
return
|
||||
}
|
||||
|
||||
userAgent := r.Header.Get("User-Agent")
|
||||
|
||||
if err := h.pushService.Subscribe(r.Context(), ownerID, req.Endpoint, req.KeyP256dh, req.KeyAuth, userAgent); err != nil {
|
||||
h.logger.Error("push subscribe failed",
|
||||
"user_id", ownerID,
|
||||
"error", err,
|
||||
)
|
||||
writeJSON(w, http.StatusInternalServerError, errorBody("server_error", "Failed to register subscription"))
|
||||
return
|
||||
}
|
||||
|
||||
h.logger.Info("push subscription registered",
|
||||
"user_id", ownerID,
|
||||
)
|
||||
|
||||
writeJSON(w, http.StatusOK, map[string]string{"status": "subscribed"})
|
||||
}
|
||||
|
||||
// Unsubscribe handles DELETE /api/push/subscribe.
|
||||
// Removes a Web Push subscription for the authenticated user.
|
||||
func (h *PushHandler) Unsubscribe(w http.ResponseWriter, r *http.Request) {
|
||||
ownerID, ok := OwnerIDFromContext(r.Context())
|
||||
if !ok {
|
||||
writeJSON(w, http.StatusUnauthorized, errorBody("unauthorized", "Authentication required"))
|
||||
return
|
||||
}
|
||||
|
||||
var req unsubscribeRequest
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
writeJSON(w, http.StatusBadRequest, errorBody("invalid_request", "Invalid JSON body"))
|
||||
return
|
||||
}
|
||||
|
||||
if req.Endpoint == "" {
|
||||
writeJSON(w, http.StatusBadRequest, errorBody("validation_error", "endpoint is required"))
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.pushService.Unsubscribe(r.Context(), ownerID, req.Endpoint); err != nil {
|
||||
h.logger.Error("push unsubscribe failed",
|
||||
"error", err,
|
||||
)
|
||||
writeJSON(w, http.StatusInternalServerError, errorBody("server_error", "Failed to remove subscription"))
|
||||
return
|
||||
}
|
||||
|
||||
writeJSON(w, http.StatusOK, map[string]string{"status": "unsubscribed"})
|
||||
}
|
||||
|
||||
// VAPIDKey handles GET /api/push/vapid-key.
|
||||
// Returns the VAPID public key needed by clients to subscribe to push notifications.
|
||||
func (h *PushHandler) VAPIDKey(w http.ResponseWriter, r *http.Request) {
|
||||
writeJSON(w, http.StatusOK, map[string]string{
|
||||
"vapid_public_key": h.pushService.GetVAPIDPublicKey(),
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,233 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"github.com/go-chi/chi/v5"
|
||||
|
||||
"github.com/synapbus/synapbus/internal/agents"
|
||||
"github.com/synapbus/synapbus/internal/auth"
|
||||
"github.com/synapbus/synapbus/internal/messaging"
|
||||
"github.com/synapbus/synapbus/internal/reactions"
|
||||
)
|
||||
|
||||
// ReactionsHandler handles REST API requests for message reactions.
|
||||
type ReactionsHandler struct {
|
||||
reactionService *reactions.Service
|
||||
msgService *messaging.MessagingService
|
||||
agentService *agents.AgentService
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
// NewReactionsHandler creates a new reactions handler.
|
||||
func NewReactionsHandler(reactionService *reactions.Service, msgService *messaging.MessagingService, agentService *agents.AgentService) *ReactionsHandler {
|
||||
return &ReactionsHandler{
|
||||
reactionService: reactionService,
|
||||
msgService: msgService,
|
||||
agentService: agentService,
|
||||
logger: slog.Default().With("component", "api.reactions"),
|
||||
}
|
||||
}
|
||||
|
||||
// Toggle handles POST /api/messages/{id}/reactions.
|
||||
func (h *ReactionsHandler) Toggle(w http.ResponseWriter, r *http.Request) {
|
||||
ownerID, ok := OwnerIDFromContext(r.Context())
|
||||
if !ok {
|
||||
writeJSON(w, http.StatusUnauthorized, errorBody("unauthorized", "Authentication required"))
|
||||
return
|
||||
}
|
||||
|
||||
id, err := strconv.ParseInt(chi.URLParam(r, "id"), 10, 64)
|
||||
if err != nil {
|
||||
writeJSON(w, http.StatusBadRequest, errorBody("invalid_id", "Invalid message ID"))
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
Reaction string `json:"reaction"`
|
||||
Metadata json.RawMessage `json:"metadata,omitempty"`
|
||||
}
|
||||
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
writeJSON(w, http.StatusBadRequest, errorBody("invalid_request", "Invalid JSON body"))
|
||||
return
|
||||
}
|
||||
|
||||
if req.Reaction == "" {
|
||||
writeJSON(w, http.StatusBadRequest, errorBody("validation_error", "Reaction type is required"))
|
||||
return
|
||||
}
|
||||
|
||||
// Verify the message exists
|
||||
msg, err := h.msgService.GetMessageByID(r.Context(), id)
|
||||
if err != nil {
|
||||
writeJSON(w, http.StatusNotFound, errorBody("not_found", "Message not found"))
|
||||
return
|
||||
}
|
||||
|
||||
// Determine the acting agent name from the session
|
||||
agentName, err := h.resolveAgentName(r, ownerID, msg)
|
||||
if err != nil {
|
||||
h.logger.Error("resolve agent name failed", "error", err)
|
||||
writeJSON(w, http.StatusBadRequest, errorBody("no_agent", err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
result, err := h.reactionService.Toggle(r.Context(), id, agentName, req.Reaction, req.Metadata)
|
||||
if err != nil {
|
||||
if err == reactions.ErrInvalidReaction {
|
||||
writeJSON(w, http.StatusBadRequest, errorBody("invalid_reaction", err.Error()))
|
||||
return
|
||||
}
|
||||
if err == reactions.ErrReactionLimit {
|
||||
writeJSON(w, http.StatusBadRequest, errorBody("reaction_limit", err.Error()))
|
||||
return
|
||||
}
|
||||
h.logger.Error("toggle reaction failed", "error", err)
|
||||
writeJSON(w, http.StatusInternalServerError, errorBody("server_error", "Failed to toggle reaction"))
|
||||
return
|
||||
}
|
||||
|
||||
// Reload reactions and workflow state for the response
|
||||
rxs, state, err := h.reactionService.GetReactions(r.Context(), id)
|
||||
if err != nil {
|
||||
h.logger.Error("get reactions failed", "error", err)
|
||||
writeJSON(w, http.StatusInternalServerError, errorBody("server_error", "Failed to get reactions"))
|
||||
return
|
||||
}
|
||||
|
||||
writeJSON(w, http.StatusOK, map[string]any{
|
||||
"action": result.Action,
|
||||
"reaction": result.Reaction,
|
||||
"reactions": rxs,
|
||||
"workflow_state": state,
|
||||
})
|
||||
}
|
||||
|
||||
// GetReactions handles GET /api/messages/{id}/reactions.
|
||||
func (h *ReactionsHandler) GetReactions(w http.ResponseWriter, r *http.Request) {
|
||||
_, ok := OwnerIDFromContext(r.Context())
|
||||
if !ok {
|
||||
writeJSON(w, http.StatusUnauthorized, errorBody("unauthorized", "Authentication required"))
|
||||
return
|
||||
}
|
||||
|
||||
id, err := strconv.ParseInt(chi.URLParam(r, "id"), 10, 64)
|
||||
if err != nil {
|
||||
writeJSON(w, http.StatusBadRequest, errorBody("invalid_id", "Invalid message ID"))
|
||||
return
|
||||
}
|
||||
|
||||
// Verify the message exists
|
||||
_, err = h.msgService.GetMessageByID(r.Context(), id)
|
||||
if err != nil {
|
||||
writeJSON(w, http.StatusNotFound, errorBody("not_found", "Message not found"))
|
||||
return
|
||||
}
|
||||
|
||||
rxs, state, err := h.reactionService.GetReactions(r.Context(), id)
|
||||
if err != nil {
|
||||
h.logger.Error("get reactions failed", "error", err)
|
||||
writeJSON(w, http.StatusInternalServerError, errorBody("server_error", "Failed to get reactions"))
|
||||
return
|
||||
}
|
||||
|
||||
writeJSON(w, http.StatusOK, map[string]any{
|
||||
"reactions": rxs,
|
||||
"workflow_state": state,
|
||||
})
|
||||
}
|
||||
|
||||
// Remove handles DELETE /api/messages/{id}/reactions/{reaction}.
|
||||
func (h *ReactionsHandler) Remove(w http.ResponseWriter, r *http.Request) {
|
||||
ownerID, ok := OwnerIDFromContext(r.Context())
|
||||
if !ok {
|
||||
writeJSON(w, http.StatusUnauthorized, errorBody("unauthorized", "Authentication required"))
|
||||
return
|
||||
}
|
||||
|
||||
id, err := strconv.ParseInt(chi.URLParam(r, "id"), 10, 64)
|
||||
if err != nil {
|
||||
writeJSON(w, http.StatusBadRequest, errorBody("invalid_id", "Invalid message ID"))
|
||||
return
|
||||
}
|
||||
|
||||
reactionType := chi.URLParam(r, "reaction")
|
||||
if reactionType == "" {
|
||||
writeJSON(w, http.StatusBadRequest, errorBody("validation_error", "Reaction type is required"))
|
||||
return
|
||||
}
|
||||
|
||||
// Verify the message exists
|
||||
msg, err := h.msgService.GetMessageByID(r.Context(), id)
|
||||
if err != nil {
|
||||
writeJSON(w, http.StatusNotFound, errorBody("not_found", "Message not found"))
|
||||
return
|
||||
}
|
||||
|
||||
// Determine the acting agent name
|
||||
agentName, err := h.resolveAgentName(r, ownerID, msg)
|
||||
if err != nil {
|
||||
h.logger.Error("resolve agent name failed", "error", err)
|
||||
writeJSON(w, http.StatusBadRequest, errorBody("no_agent", err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.reactionService.Remove(r.Context(), id, agentName, reactionType); err != nil {
|
||||
if err == reactions.ErrInvalidReaction {
|
||||
writeJSON(w, http.StatusBadRequest, errorBody("invalid_reaction", err.Error()))
|
||||
return
|
||||
}
|
||||
h.logger.Error("remove reaction failed", "error", err)
|
||||
writeJSON(w, http.StatusInternalServerError, errorBody("server_error", "Failed to remove reaction"))
|
||||
return
|
||||
}
|
||||
|
||||
// Reload reactions and workflow state for the response
|
||||
rxs, state, err := h.reactionService.GetReactions(r.Context(), id)
|
||||
if err != nil {
|
||||
h.logger.Error("get reactions failed", "error", err)
|
||||
writeJSON(w, http.StatusInternalServerError, errorBody("server_error", "Failed to get reactions"))
|
||||
return
|
||||
}
|
||||
|
||||
writeJSON(w, http.StatusOK, map[string]any{
|
||||
"status": "removed",
|
||||
"reactions": rxs,
|
||||
"workflow_state": state,
|
||||
})
|
||||
}
|
||||
|
||||
// resolveAgentName determines the agent name for the current session user.
|
||||
// For session-authenticated users (Web UI), it returns the human agent.
|
||||
// For API key / OAuth, it falls back to the first owned agent.
|
||||
func (h *ReactionsHandler) resolveAgentName(r *http.Request, ownerID int64, msg *messaging.Message) (string, error) {
|
||||
if _, isSession := auth.SessionIDFromContext(r.Context()); isSession {
|
||||
humanAgent, err := h.agentService.GetHumanAgentForUser(r.Context(), ownerID)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if humanAgent == nil {
|
||||
return "", fmt.Errorf("no human agent found for this user")
|
||||
}
|
||||
return humanAgent.Name, nil
|
||||
}
|
||||
|
||||
// Non-session: find an owned agent
|
||||
ownedAgents, err := h.agentService.ListAgents(r.Context(), ownerID)
|
||||
if err != nil || len(ownedAgents) == 0 {
|
||||
return "", fmt.Errorf("no agents registered")
|
||||
}
|
||||
|
||||
// Prefer human-type agent
|
||||
for _, a := range ownedAgents {
|
||||
if a.Type == "human" {
|
||||
return a.Name, nil
|
||||
}
|
||||
}
|
||||
return ownedAgents[0].Name, nil
|
||||
}
|
||||
+61
-3
@@ -1,6 +1,7 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"net/http"
|
||||
|
||||
"github.com/go-chi/chi/v5"
|
||||
@@ -11,6 +12,8 @@ import (
|
||||
"github.com/synapbus/synapbus/internal/channels"
|
||||
"github.com/synapbus/synapbus/internal/k8s"
|
||||
"github.com/synapbus/synapbus/internal/messaging"
|
||||
"github.com/synapbus/synapbus/internal/push"
|
||||
"github.com/synapbus/synapbus/internal/reactions"
|
||||
"github.com/synapbus/synapbus/internal/trace"
|
||||
"github.com/synapbus/synapbus/internal/webhooks"
|
||||
)
|
||||
@@ -30,8 +33,13 @@ type RouterConfig struct {
|
||||
WebhookStore webhooks.WebhookStore
|
||||
K8sService *k8s.K8sService
|
||||
K8sStore k8s.K8sStore
|
||||
ReactionService *reactions.Service
|
||||
PushService *push.Service
|
||||
SSEHub *SSEHub
|
||||
Broadcaster *SSEBroadcaster
|
||||
SessionMiddleware func(http.Handler) http.Handler
|
||||
DB *sql.DB
|
||||
Version string
|
||||
}
|
||||
|
||||
// NewRouter creates a chi router with all API routes configured.
|
||||
@@ -88,9 +96,10 @@ func NewRouterWithConfig(cfg RouterConfig) chi.Router {
|
||||
notificationsHandler := NewNotificationsHandler(cfg.MsgService, cfg.AgentService, cfg.ChannelService)
|
||||
|
||||
// Wire up SSE broadcaster for real-time events
|
||||
if cfg.SSEHub != nil {
|
||||
broadcaster := NewSSEBroadcaster(cfg.SSEHub, cfg.AgentService, cfg.ChannelService)
|
||||
messagesHandler.SetBroadcaster(broadcaster)
|
||||
if cfg.Broadcaster != nil {
|
||||
messagesHandler.SetBroadcaster(cfg.Broadcaster)
|
||||
} else if cfg.SSEHub != nil {
|
||||
messagesHandler.SetBroadcaster(NewSSEBroadcaster(cfg.SSEHub, cfg.AgentService, cfg.ChannelService))
|
||||
}
|
||||
|
||||
r.Group(func(r chi.Router) {
|
||||
@@ -135,9 +144,24 @@ func NewRouterWithConfig(cfg RouterConfig) chi.Router {
|
||||
})
|
||||
}
|
||||
|
||||
// Reactions
|
||||
if cfg.ReactionService != nil {
|
||||
reactionsHandler := NewReactionsHandler(cfg.ReactionService, cfg.MsgService, cfg.AgentService)
|
||||
r.Group(func(r chi.Router) {
|
||||
r.Use(authMiddleware)
|
||||
|
||||
r.Post("/api/messages/{id}/reactions", reactionsHandler.Toggle)
|
||||
r.Get("/api/messages/{id}/reactions", reactionsHandler.GetReactions)
|
||||
r.Delete("/api/messages/{id}/reactions/{reaction}", reactionsHandler.Remove)
|
||||
})
|
||||
}
|
||||
|
||||
// Channels
|
||||
if cfg.ChannelService != nil {
|
||||
channelsHandler := NewChannelsHandler(cfg.ChannelService, cfg.AgentService, cfg.MsgService)
|
||||
if cfg.ReactionService != nil {
|
||||
channelsHandler.SetReactionService(cfg.ReactionService)
|
||||
}
|
||||
r.Group(func(r chi.Router) {
|
||||
r.Use(authMiddleware)
|
||||
|
||||
@@ -145,6 +169,8 @@ func NewRouterWithConfig(cfg RouterConfig) chi.Router {
|
||||
r.Get("/api/channels/{name}", channelsHandler.GetChannel)
|
||||
r.Post("/api/channels", channelsHandler.CreateChannel)
|
||||
r.Get("/api/channels/{name}/messages", channelsHandler.ChannelMessages)
|
||||
r.Get("/api/channels/{name}/messages/by-state", channelsHandler.ListByState)
|
||||
r.Put("/api/channels/{name}/settings", channelsHandler.UpdateSettings)
|
||||
r.Post("/api/channels/{name}/join", channelsHandler.JoinChannel)
|
||||
r.Post("/api/channels/{name}/leave", channelsHandler.LeaveChannel)
|
||||
})
|
||||
@@ -194,6 +220,38 @@ func NewRouterWithConfig(cfg RouterConfig) chi.Router {
|
||||
r.Get("/api/k8s/job-runs/{id}/logs", k8sHandler.JobRunLogs)
|
||||
})
|
||||
}
|
||||
|
||||
// Push Notifications
|
||||
if cfg.PushService != nil {
|
||||
pushHandler := NewPushHandler(cfg.PushService)
|
||||
// VAPID key endpoint is unauthenticated (needed before subscription)
|
||||
r.Get("/api/push/vapid-key", pushHandler.VAPIDKey)
|
||||
r.Group(func(r chi.Router) {
|
||||
r.Use(authMiddleware)
|
||||
|
||||
r.Post("/api/push/subscribe", pushHandler.Subscribe)
|
||||
r.Delete("/api/push/subscribe", pushHandler.Unsubscribe)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Analytics (authenticated, requires DB)
|
||||
if cfg.DB != nil {
|
||||
analyticsHandler := NewAnalyticsHandler(cfg.DB, cfg.AgentService, cfg.ChannelService)
|
||||
r.Group(func(r chi.Router) {
|
||||
r.Use(authMiddleware)
|
||||
|
||||
r.Get("/api/analytics/timeline", analyticsHandler.Timeline)
|
||||
r.Get("/api/analytics/top-agents", analyticsHandler.TopAgents)
|
||||
r.Get("/api/analytics/top-channels", analyticsHandler.TopChannels)
|
||||
r.Get("/api/analytics/summary", analyticsHandler.Summary)
|
||||
})
|
||||
}
|
||||
|
||||
// Version (unauthenticated)
|
||||
if cfg.Version != "" {
|
||||
versionHandler := NewVersionHandler(cfg.Version)
|
||||
r.Get("/api/version", versionHandler.GetVersion)
|
||||
}
|
||||
|
||||
// Metrics endpoint (unauthenticated, only registered when enabled)
|
||||
|
||||
@@ -0,0 +1,24 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
)
|
||||
|
||||
// VersionHandler serves version information.
|
||||
// The version is set at build time via -ldflags.
|
||||
type VersionHandler struct {
|
||||
version string
|
||||
}
|
||||
|
||||
// NewVersionHandler creates a new version handler.
|
||||
func NewVersionHandler(version string) *VersionHandler {
|
||||
return &VersionHandler{version: version}
|
||||
}
|
||||
|
||||
// GetVersion handles GET /api/version.
|
||||
func (h *VersionHandler) GetVersion(w http.ResponseWriter, r *http.Request) {
|
||||
writeJSON(w, http.StatusOK, map[string]string{
|
||||
"version": h.version,
|
||||
"repo": "https://github.com/synapbus/synapbus",
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,87 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestGetVersion(t *testing.T) {
|
||||
handler := NewVersionHandler("v0.7.0")
|
||||
|
||||
req := httptest.NewRequest("GET", "/api/version", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
handler.GetVersion(rr, req)
|
||||
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want %d", rr.Code, http.StatusOK)
|
||||
}
|
||||
|
||||
ct := rr.Header().Get("Content-Type")
|
||||
if ct != "application/json" {
|
||||
t.Errorf("Content-Type = %q, want application/json", ct)
|
||||
}
|
||||
|
||||
var resp struct {
|
||||
Version string `json:"version"`
|
||||
Repo string `json:"repo"`
|
||||
}
|
||||
if err := json.Unmarshal(rr.Body.Bytes(), &resp); err != nil {
|
||||
t.Fatalf("unmarshal: %v", err)
|
||||
}
|
||||
|
||||
if resp.Version != "v0.7.0" {
|
||||
t.Errorf("version = %q, want v0.7.0", resp.Version)
|
||||
}
|
||||
if resp.Repo != "https://github.com/synapbus/synapbus" {
|
||||
t.Errorf("repo = %q, want https://github.com/synapbus/synapbus", resp.Repo)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetVersion_DevBuild(t *testing.T) {
|
||||
handler := NewVersionHandler("dev")
|
||||
|
||||
req := httptest.NewRequest("GET", "/api/version", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
handler.GetVersion(rr, req)
|
||||
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want %d", rr.Code, http.StatusOK)
|
||||
}
|
||||
|
||||
var resp struct {
|
||||
Version string `json:"version"`
|
||||
Repo string `json:"repo"`
|
||||
}
|
||||
json.Unmarshal(rr.Body.Bytes(), &resp)
|
||||
|
||||
if resp.Version != "dev" {
|
||||
t.Errorf("version = %q, want dev", resp.Version)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetVersion_ResponseFormat(t *testing.T) {
|
||||
handler := NewVersionHandler("v1.2.3")
|
||||
|
||||
req := httptest.NewRequest("GET", "/api/version", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
handler.GetVersion(rr, req)
|
||||
|
||||
// Verify the response is valid JSON with exactly the expected keys
|
||||
var raw map[string]any
|
||||
if err := json.Unmarshal(rr.Body.Bytes(), &raw); err != nil {
|
||||
t.Fatalf("unmarshal: %v", err)
|
||||
}
|
||||
|
||||
if len(raw) != 2 {
|
||||
t.Errorf("response has %d keys, want 2", len(raw))
|
||||
}
|
||||
|
||||
if _, ok := raw["version"]; !ok {
|
||||
t.Error("response missing 'version' key")
|
||||
}
|
||||
if _, ok := raw["repo"]; !ok {
|
||||
t.Error("response missing 'repo' key")
|
||||
}
|
||||
}
|
||||
@@ -118,3 +118,27 @@ func DefaultFilename(mimeType string) string {
|
||||
func IsImageType(mimeType string) bool {
|
||||
return imageTypes[mimeType]
|
||||
}
|
||||
|
||||
// IsAllowedType returns true if the MIME type is allowed for upload.
|
||||
// Allowed: image/*, application/pdf, text/*.
|
||||
func IsAllowedType(mimeType string) bool {
|
||||
// Normalize: strip parameters like "; charset=utf-8".
|
||||
base := mimeType
|
||||
if idx := strings.Index(mimeType, ";"); idx >= 0 {
|
||||
base = strings.TrimSpace(mimeType[:idx])
|
||||
}
|
||||
if strings.HasPrefix(base, "image/") {
|
||||
return true
|
||||
}
|
||||
if base == "application/pdf" {
|
||||
return true
|
||||
}
|
||||
if strings.HasPrefix(base, "text/") {
|
||||
return true
|
||||
}
|
||||
// Also allow JSON and XML which may be detected as application/*
|
||||
if base == "application/json" || base == "application/xml" {
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
@@ -154,6 +154,35 @@ func TestIsImageType(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsAllowedType(t *testing.T) {
|
||||
tests := []struct {
|
||||
mimeType string
|
||||
want bool
|
||||
}{
|
||||
{"image/png", true},
|
||||
{"image/jpeg", true},
|
||||
{"image/gif", true},
|
||||
{"application/pdf", true},
|
||||
{"text/plain", true},
|
||||
{"text/csv", true},
|
||||
{"text/plain; charset=utf-8", true},
|
||||
{"application/json", true},
|
||||
{"application/octet-stream", false},
|
||||
{"application/zip", false},
|
||||
{"application/x-executable", false},
|
||||
{"video/mp4", false},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.mimeType, func(t *testing.T) {
|
||||
got := IsAllowedType(tt.mimeType)
|
||||
if got != tt.want {
|
||||
t.Errorf("IsAllowedType(%q) = %v, want %v", tt.mimeType, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func min(a, b int) int {
|
||||
if a < b {
|
||||
return a
|
||||
|
||||
@@ -11,10 +11,11 @@ const MaxFileSize = 50 * 1024 * 1024 // 50 MB
|
||||
|
||||
// Sentinel errors.
|
||||
var (
|
||||
ErrNotFound = errors.New("attachment not found")
|
||||
ErrFileTooLarge = errors.New("file exceeds maximum size of 50MB")
|
||||
ErrEmptyFile = errors.New("empty file not allowed")
|
||||
ErrFileMissing = errors.New("attachment file missing from disk")
|
||||
ErrNotFound = errors.New("attachment not found")
|
||||
ErrFileTooLarge = errors.New("file exceeds maximum size of 50MB")
|
||||
ErrEmptyFile = errors.New("empty file not allowed")
|
||||
ErrFileMissing = errors.New("attachment file missing from disk")
|
||||
ErrUnsupportedType = errors.New("unsupported file type: only images (jpg, png, gif, webp, svg), PDFs, and text files are allowed")
|
||||
)
|
||||
|
||||
// Attachment represents the metadata for a stored file.
|
||||
|
||||
@@ -60,6 +60,11 @@ func (s *Service) Upload(ctx context.Context, req UploadRequest) (*UploadResult,
|
||||
mimeType = DetectMIMEType(sniffBuf, req.Filename)
|
||||
}
|
||||
|
||||
// Validate file type against allowlist.
|
||||
if !IsAllowedType(mimeType) {
|
||||
return nil, ErrUnsupportedType
|
||||
}
|
||||
|
||||
// Assign default filename if missing.
|
||||
filename := req.Filename
|
||||
if filename == "" {
|
||||
|
||||
@@ -202,6 +202,80 @@ func TestService_Dedup(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestService_Upload_FileTypeValidation(t *testing.T) {
|
||||
// PNG magic bytes.
|
||||
pngContent := []byte{0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a, 0x00, 0x00}
|
||||
// PDF magic bytes.
|
||||
pdfContent := []byte("%PDF-1.4 some pdf content here")
|
||||
// Plain text content.
|
||||
textContent := []byte("just some plain text content")
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
content []byte
|
||||
filename string
|
||||
mimeType string
|
||||
wantErr error
|
||||
}{
|
||||
{
|
||||
name: "valid image upload",
|
||||
content: pngContent,
|
||||
filename: "photo.png",
|
||||
wantErr: nil,
|
||||
},
|
||||
{
|
||||
name: "valid PDF upload",
|
||||
content: pdfContent,
|
||||
filename: "report.pdf",
|
||||
wantErr: nil,
|
||||
},
|
||||
{
|
||||
name: "valid text file upload",
|
||||
content: textContent,
|
||||
filename: "notes.txt",
|
||||
wantErr: nil,
|
||||
},
|
||||
{
|
||||
name: "invalid type zip rejected",
|
||||
content: []byte("not real zip content"),
|
||||
filename: "archive.zip",
|
||||
mimeType: "application/zip",
|
||||
wantErr: ErrUnsupportedType,
|
||||
},
|
||||
{
|
||||
name: "invalid type executable rejected",
|
||||
content: []byte{0x7f, 0x45, 0x4c, 0x46},
|
||||
filename: "program.exe",
|
||||
mimeType: "application/x-executable",
|
||||
wantErr: ErrUnsupportedType,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
svc, _ := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
_, err := svc.Upload(ctx, UploadRequest{
|
||||
Content: bytes.NewReader(tt.content),
|
||||
Filename: tt.filename,
|
||||
MIMEType: tt.mimeType,
|
||||
UploadedBy: "agent-a",
|
||||
})
|
||||
|
||||
if tt.wantErr != nil {
|
||||
if err != tt.wantErr {
|
||||
t.Errorf("expected error %v, got %v", tt.wantErr, err)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("Upload: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestService_GarbageCollect(t *testing.T) {
|
||||
svc, db := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
@@ -227,6 +227,59 @@ func (h *Handlers) HandleMe(w http.ResponseWriter, r *http.Request) {
|
||||
json.NewEncoder(w).Encode(user)
|
||||
}
|
||||
|
||||
// HandleUpdateProfile handles PUT /api/auth/profile.
|
||||
func (h *Handlers) HandleUpdateProfile(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPut {
|
||||
writeError(w, http.StatusMethodNotAllowed, "method_not_allowed", "PUT required")
|
||||
return
|
||||
}
|
||||
|
||||
user, ok := UserFromContext(r.Context())
|
||||
if !ok {
|
||||
writeError(w, http.StatusUnauthorized, "unauthorized", "Not authenticated")
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
DisplayName string `json:"display_name"`
|
||||
}
|
||||
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
writeError(w, http.StatusBadRequest, "invalid_request", "Invalid JSON body")
|
||||
return
|
||||
}
|
||||
|
||||
if strings.TrimSpace(req.DisplayName) == "" {
|
||||
writeError(w, http.StatusBadRequest, "invalid_request", "Display name cannot be empty")
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.userStore.UpdateDisplayName(r.Context(), user.ID, req.DisplayName); err != nil {
|
||||
h.logger.Error("update display name failed", "error", err)
|
||||
writeError(w, http.StatusInternalServerError, "server_error", "Failed to update display name")
|
||||
return
|
||||
}
|
||||
|
||||
// Fetch updated user
|
||||
updated, err := h.userStore.GetUserByID(r.Context(), user.ID)
|
||||
if err != nil {
|
||||
h.logger.Error("fetch updated user failed", "error", err)
|
||||
writeError(w, http.StatusInternalServerError, "server_error", "Profile updated but failed to fetch result")
|
||||
return
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(map[string]any{
|
||||
"message": "Profile updated",
|
||||
"user": map[string]any{
|
||||
"id": updated.ID,
|
||||
"username": updated.Username,
|
||||
"display_name": updated.DisplayName,
|
||||
"role": updated.Role,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// HandleChangePassword handles PUT /auth/password.
|
||||
func (h *Handlers) HandleChangePassword(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPut {
|
||||
@@ -295,13 +348,21 @@ func (h *Handlers) HandleOAuthMetadata(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
baseURL := h.config.IssuerURL
|
||||
if baseURL == "" {
|
||||
// Fall back to request Host header
|
||||
if baseURL == "" || baseURL == "auto" {
|
||||
// Derive from request headers (supports reverse proxies and tunnels)
|
||||
scheme := "http"
|
||||
if r.TLS != nil {
|
||||
scheme = "https"
|
||||
}
|
||||
baseURL = fmt.Sprintf("%s://%s", scheme, r.Host)
|
||||
// Cloudflare, nginx, and other proxies set these headers
|
||||
if proto := r.Header.Get("X-Forwarded-Proto"); proto != "" {
|
||||
scheme = proto
|
||||
}
|
||||
host := r.Host
|
||||
if fwdHost := r.Header.Get("X-Forwarded-Host"); fwdHost != "" {
|
||||
host = fwdHost
|
||||
}
|
||||
baseURL = fmt.Sprintf("%s://%s", scheme, host)
|
||||
}
|
||||
|
||||
metadata := map[string]any{
|
||||
@@ -816,13 +877,28 @@ const authorizeTemplateHTML = `<!DOCTYPE html>
|
||||
<body>
|
||||
<div class="card">
|
||||
<div class="logo">
|
||||
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 24 24" fill="none" stroke-width="1.8">
|
||||
<path d="M12 2L2 7l10 5 10-5-10-5z" stroke="#36c5f0"/>
|
||||
<path d="M2 17l10 5 10-5" stroke="#7c3aed"/>
|
||||
<path d="M2 12l10 5 10-5" stroke="url(#g)"/>
|
||||
<defs><linearGradient id="g" x1="2" y1="12" x2="22" y2="17" gradientUnits="userSpaceOnUse">
|
||||
<stop stop-color="#36c5f0"/><stop offset="1" stop-color="#7c3aed"/>
|
||||
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 32 32">
|
||||
<defs><linearGradient id="bg" x1="0" y1="0" x2="32" y2="32" gradientUnits="userSpaceOnUse">
|
||||
<stop stop-color="#7c3aed"/><stop offset="1" stop-color="#06b6d4"/>
|
||||
</linearGradient></defs>
|
||||
<rect width="32" height="32" rx="7" fill="url(#bg)"/>
|
||||
<g transform="translate(4,4)">
|
||||
<circle cx="12" cy="5" r="1.8" fill="#c4b5fd"/>
|
||||
<circle cx="6" cy="10" r="1.5" fill="#a78bfa"/>
|
||||
<circle cx="18" cy="9" r="1.5" fill="#c4b5fd"/>
|
||||
<circle cx="12" cy="13" r="2" fill="#67e8f9"/>
|
||||
<circle cx="5" cy="17" r="1.5" fill="#a78bfa"/>
|
||||
<circle cx="19" cy="17" r="1.3" fill="#a78bfa"/>
|
||||
<circle cx="11" cy="20" r="1.3" fill="#c4b5fd"/>
|
||||
<line x1="12" y1="5" x2="6" y2="10" stroke="white" stroke-width="0.5" opacity="0.5"/>
|
||||
<line x1="12" y1="5" x2="18" y2="9" stroke="white" stroke-width="0.5" opacity="0.5"/>
|
||||
<line x1="6" y1="10" x2="12" y2="13" stroke="white" stroke-width="0.5" opacity="0.5"/>
|
||||
<line x1="18" y1="9" x2="12" y2="13" stroke="white" stroke-width="0.5" opacity="0.5"/>
|
||||
<line x1="12" y1="13" x2="5" y2="17" stroke="white" stroke-width="0.5" opacity="0.5"/>
|
||||
<line x1="12" y1="13" x2="19" y2="17" stroke="white" stroke-width="0.5" opacity="0.5"/>
|
||||
<line x1="5" y1="17" x2="11" y2="20" stroke="white" stroke-width="0.5" opacity="0.3"/>
|
||||
<line x1="6" y1="10" x2="5" y2="17" stroke="white" stroke-width="0.5" opacity="0.3"/>
|
||||
</g>
|
||||
</svg>
|
||||
<h1>SynapBus</h1>
|
||||
<p>Agent Authorization</p>
|
||||
|
||||
@@ -0,0 +1,93 @@
|
||||
package idp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
"os"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// LoadConfig reads environment variables and returns a slice of enabled identity providers.
|
||||
// Providers are only included if both client ID and client secret are set.
|
||||
func LoadConfig(baseURL string) []Provider {
|
||||
var providers []Provider
|
||||
|
||||
// GitHub
|
||||
if clientID, clientSecret := os.Getenv("SYNAPBUS_IDP_GITHUB_CLIENT_ID"), os.Getenv("SYNAPBUS_IDP_GITHUB_CLIENT_SECRET"); clientID != "" && clientSecret != "" {
|
||||
redirectURL := baseURL + "/auth/callback/github"
|
||||
providers = append(providers, NewGitHubProvider(clientID, clientSecret, redirectURL))
|
||||
slog.Info("IdP enabled", "provider", "github")
|
||||
}
|
||||
|
||||
// Google (OIDC)
|
||||
if clientID, clientSecret := os.Getenv("SYNAPBUS_IDP_GOOGLE_CLIENT_ID"), os.Getenv("SYNAPBUS_IDP_GOOGLE_CLIENT_SECRET"); clientID != "" && clientSecret != "" {
|
||||
var allowedDomains []string
|
||||
if domains := os.Getenv("SYNAPBUS_IDP_GOOGLE_ALLOWED_DOMAINS"); domains != "" {
|
||||
for _, d := range strings.Split(domains, ",") {
|
||||
d = strings.TrimSpace(d)
|
||||
if d != "" {
|
||||
allowedDomains = append(allowedDomains, d)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
p, err := NewOIDCProvider(ctx, OIDCConfig{
|
||||
ID: "google",
|
||||
DisplayName: "Google",
|
||||
IssuerURL: "https://accounts.google.com",
|
||||
ClientID: clientID,
|
||||
ClientSecret: clientSecret,
|
||||
RedirectURL: baseURL + "/auth/callback/google",
|
||||
AllowedDomains: allowedDomains,
|
||||
})
|
||||
if err != nil {
|
||||
slog.Error("failed to initialize Google OIDC provider", "error", err)
|
||||
} else {
|
||||
providers = append(providers, p)
|
||||
slog.Info("IdP enabled", "provider", "google", "allowed_domains", allowedDomains)
|
||||
}
|
||||
}
|
||||
|
||||
// Azure AD (OIDC)
|
||||
if clientID, clientSecret := os.Getenv("SYNAPBUS_IDP_AZUREAD_CLIENT_ID"), os.Getenv("SYNAPBUS_IDP_AZUREAD_CLIENT_SECRET"); clientID != "" && clientSecret != "" {
|
||||
tenantID := os.Getenv("SYNAPBUS_IDP_AZUREAD_TENANT_ID")
|
||||
if tenantID == "" {
|
||||
tenantID = "common" // multi-tenant by default
|
||||
}
|
||||
|
||||
var groupMapping map[string]string
|
||||
if gm := os.Getenv("SYNAPBUS_IDP_AZUREAD_GROUP_MAPPING"); gm != "" {
|
||||
if err := json.Unmarshal([]byte(gm), &groupMapping); err != nil {
|
||||
slog.Error("failed to parse Azure AD group mapping", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
issuerURL := "https://login.microsoftonline.com/" + tenantID + "/v2.0"
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
p, err := NewOIDCProvider(ctx, OIDCConfig{
|
||||
ID: "azuread",
|
||||
DisplayName: "Microsoft",
|
||||
IssuerURL: issuerURL,
|
||||
ClientID: clientID,
|
||||
ClientSecret: clientSecret,
|
||||
RedirectURL: baseURL + "/auth/callback/azuread",
|
||||
GroupMapping: groupMapping,
|
||||
})
|
||||
if err != nil {
|
||||
slog.Error("failed to initialize Azure AD OIDC provider", "error", err)
|
||||
} else {
|
||||
providers = append(providers, p)
|
||||
slog.Info("IdP enabled", "provider", "azuread", "tenant_id", tenantID)
|
||||
}
|
||||
}
|
||||
|
||||
return providers
|
||||
}
|
||||
@@ -0,0 +1,140 @@
|
||||
package idp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"golang.org/x/oauth2"
|
||||
"golang.org/x/oauth2/github"
|
||||
)
|
||||
|
||||
// GitHubProvider implements GitHub OAuth authentication.
|
||||
// GitHub does not support OIDC discovery, so this uses plain OAuth 2.0
|
||||
// with GitHub's user API endpoints.
|
||||
type GitHubProvider struct {
|
||||
config *oauth2.Config
|
||||
}
|
||||
|
||||
// NewGitHubProvider creates a new GitHub OAuth provider.
|
||||
func NewGitHubProvider(clientID, clientSecret, redirectURL string) *GitHubProvider {
|
||||
return &GitHubProvider{
|
||||
config: &oauth2.Config{
|
||||
ClientID: clientID,
|
||||
ClientSecret: clientSecret,
|
||||
RedirectURL: redirectURL,
|
||||
Scopes: []string{"read:user", "user:email"},
|
||||
Endpoint: github.Endpoint,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (p *GitHubProvider) ID() string { return "github" }
|
||||
func (p *GitHubProvider) Type() string { return "oauth" }
|
||||
func (p *GitHubProvider) DisplayName() string { return "GitHub" }
|
||||
|
||||
func (p *GitHubProvider) AuthCodeURL(state string) string {
|
||||
return p.config.AuthCodeURL(state)
|
||||
}
|
||||
|
||||
func (p *GitHubProvider) Exchange(ctx context.Context, code string) (*ExternalUser, error) {
|
||||
token, err := p.config.Exchange(ctx, code)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("github: exchange code: %w", err)
|
||||
}
|
||||
|
||||
client := p.config.Client(ctx, token)
|
||||
|
||||
// Fetch user profile
|
||||
userInfo, err := fetchGitHubUser(client)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Fetch primary verified email
|
||||
email, err := fetchGitHubPrimaryEmail(client)
|
||||
if err != nil {
|
||||
// Non-fatal: email might not be available
|
||||
email = ""
|
||||
}
|
||||
|
||||
// If user profile has an email and we didn't get one from the emails endpoint, use it
|
||||
if email == "" {
|
||||
if e, ok := userInfo["email"].(string); ok && e != "" {
|
||||
email = e
|
||||
}
|
||||
}
|
||||
|
||||
idNum, _ := userInfo["id"].(float64)
|
||||
login, _ := userInfo["login"].(string)
|
||||
name, _ := userInfo["name"].(string)
|
||||
|
||||
return &ExternalUser{
|
||||
ProviderID: "github",
|
||||
ExternalID: strconv.FormatInt(int64(idNum), 10),
|
||||
Email: email,
|
||||
Name: name,
|
||||
Username: login,
|
||||
RawClaims: userInfo,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func fetchGitHubUser(client *http.Client) (map[string]any, error) {
|
||||
resp, err := client.Get("https://api.github.com/user")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("github: fetch user: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
return nil, fmt.Errorf("github: user API returned %d: %s", resp.StatusCode, body)
|
||||
}
|
||||
|
||||
var user map[string]any
|
||||
if err := json.NewDecoder(resp.Body).Decode(&user); err != nil {
|
||||
return nil, fmt.Errorf("github: decode user: %w", err)
|
||||
}
|
||||
return user, nil
|
||||
}
|
||||
|
||||
// githubEmail represents an email entry from the GitHub emails API.
|
||||
type githubEmail struct {
|
||||
Email string `json:"email"`
|
||||
Primary bool `json:"primary"`
|
||||
Verified bool `json:"verified"`
|
||||
}
|
||||
|
||||
func fetchGitHubPrimaryEmail(client *http.Client) (string, error) {
|
||||
resp, err := client.Get("https://api.github.com/user/emails")
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("github: fetch emails: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return "", fmt.Errorf("github: emails API returned %d", resp.StatusCode)
|
||||
}
|
||||
|
||||
var emails []githubEmail
|
||||
if err := json.NewDecoder(resp.Body).Decode(&emails); err != nil {
|
||||
return "", fmt.Errorf("github: decode emails: %w", err)
|
||||
}
|
||||
|
||||
// Find primary verified email
|
||||
for _, e := range emails {
|
||||
if e.Primary && e.Verified {
|
||||
return e.Email, nil
|
||||
}
|
||||
}
|
||||
// Fallback: first verified email
|
||||
for _, e := range emails {
|
||||
if e.Verified {
|
||||
return e.Email, nil
|
||||
}
|
||||
}
|
||||
return "", fmt.Errorf("github: no verified email found")
|
||||
}
|
||||
@@ -0,0 +1,349 @@
|
||||
package idp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"database/sql"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/go-chi/chi/v5"
|
||||
"github.com/synapbus/synapbus/internal/auth"
|
||||
)
|
||||
|
||||
// AgentProvisioner creates a human agent and personal channel after IdP login.
|
||||
// This is a narrow interface to avoid importing the agents/channels packages.
|
||||
type AgentProvisioner interface {
|
||||
ProvisionHumanAgent(ctx context.Context, username, displayName string, ownerID int64) error
|
||||
}
|
||||
|
||||
// Handlers holds HTTP handlers for external identity provider authentication.
|
||||
type Handlers struct {
|
||||
providers map[string]Provider
|
||||
providerList []Provider // preserve order for listing
|
||||
idStore *UserIdentityStore
|
||||
userStore auth.UserStore
|
||||
sessionStore auth.SessionStore
|
||||
agentProvisioner AgentProvisioner
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
// NewHandlers creates a new set of IdP HTTP handlers.
|
||||
func NewHandlers(
|
||||
providers []Provider,
|
||||
idStore *UserIdentityStore,
|
||||
userStore auth.UserStore,
|
||||
sessionStore auth.SessionStore,
|
||||
agentProvisioner AgentProvisioner,
|
||||
) *Handlers {
|
||||
providerMap := make(map[string]Provider, len(providers))
|
||||
for _, p := range providers {
|
||||
providerMap[p.ID()] = p
|
||||
}
|
||||
return &Handlers{
|
||||
providers: providerMap,
|
||||
providerList: providers,
|
||||
idStore: idStore,
|
||||
userStore: userStore,
|
||||
sessionStore: sessionStore,
|
||||
agentProvisioner: agentProvisioner,
|
||||
logger: slog.Default().With("component", "idp"),
|
||||
}
|
||||
}
|
||||
|
||||
// HandleListProviders returns the list of enabled identity providers.
|
||||
// GET /auth/providers
|
||||
func (h *Handlers) HandleListProviders(w http.ResponseWriter, r *http.Request) {
|
||||
providers := make([]ProviderInfo, 0, len(h.providerList))
|
||||
for _, p := range h.providerList {
|
||||
providers = append(providers, ProviderInfo{
|
||||
ID: p.ID(),
|
||||
Type: p.Type(),
|
||||
DisplayName: p.DisplayName(),
|
||||
})
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(map[string]any{
|
||||
"providers": providers,
|
||||
})
|
||||
}
|
||||
|
||||
// HandleLogin initiates the OAuth/OIDC flow by redirecting to the IdP.
|
||||
// GET /auth/login/{provider}
|
||||
func (h *Handlers) HandleLogin(w http.ResponseWriter, r *http.Request) {
|
||||
providerID := chi.URLParam(r, "provider")
|
||||
provider, ok := h.providers[providerID]
|
||||
if !ok {
|
||||
http.Error(w, fmt.Sprintf(`{"error":"unknown_provider","message":"Provider %q not found"}`, providerID), http.StatusNotFound)
|
||||
return
|
||||
}
|
||||
|
||||
state, err := generateState()
|
||||
if err != nil {
|
||||
h.logger.Error("failed to generate state", "error", err)
|
||||
http.Error(w, `{"error":"server_error","message":"Failed to generate state"}`, http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
// Store state in a short-lived cookie for CSRF protection
|
||||
http.SetCookie(w, &http.Cookie{
|
||||
Name: "idp_state",
|
||||
Value: state,
|
||||
Path: "/auth/callback/" + providerID,
|
||||
HttpOnly: true,
|
||||
SameSite: http.SameSiteLaxMode,
|
||||
MaxAge: 600, // 10 minutes
|
||||
})
|
||||
|
||||
authURL := provider.AuthCodeURL(state)
|
||||
http.Redirect(w, r, authURL, http.StatusTemporaryRedirect)
|
||||
}
|
||||
|
||||
// HandleCallback processes the OAuth/OIDC callback from the IdP.
|
||||
// GET /auth/callback/{provider}
|
||||
func (h *Handlers) HandleCallback(w http.ResponseWriter, r *http.Request) {
|
||||
providerID := chi.URLParam(r, "provider")
|
||||
provider, ok := h.providers[providerID]
|
||||
if !ok {
|
||||
http.Error(w, "Unknown provider", http.StatusNotFound)
|
||||
return
|
||||
}
|
||||
|
||||
// Verify state parameter
|
||||
stateCookie, err := r.Cookie("idp_state")
|
||||
if err != nil || stateCookie.Value == "" {
|
||||
http.Error(w, "Missing state cookie", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
if r.URL.Query().Get("state") != stateCookie.Value {
|
||||
http.Error(w, "State mismatch — possible CSRF attack", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
// Clear state cookie
|
||||
http.SetCookie(w, &http.Cookie{
|
||||
Name: "idp_state",
|
||||
Value: "",
|
||||
Path: "/auth/callback/" + providerID,
|
||||
HttpOnly: true,
|
||||
MaxAge: -1,
|
||||
})
|
||||
|
||||
// Check for error from IdP
|
||||
if errParam := r.URL.Query().Get("error"); errParam != "" {
|
||||
desc := r.URL.Query().Get("error_description")
|
||||
h.logger.Warn("IdP returned error", "provider", providerID, "error", errParam, "description", desc)
|
||||
http.Redirect(w, r, "/login?error="+errParam, http.StatusTemporaryRedirect)
|
||||
return
|
||||
}
|
||||
|
||||
code := r.URL.Query().Get("code")
|
||||
if code == "" {
|
||||
http.Error(w, "Missing authorization code", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
// Exchange code for user info
|
||||
extUser, err := provider.Exchange(r.Context(), code)
|
||||
if err != nil {
|
||||
h.logger.Error("IdP exchange failed", "provider", providerID, "error", err)
|
||||
http.Redirect(w, r, "/login?error=exchange_failed", http.StatusTemporaryRedirect)
|
||||
return
|
||||
}
|
||||
|
||||
// Find or create local user
|
||||
user, err := h.findOrCreateUser(r.Context(), extUser)
|
||||
if err != nil {
|
||||
h.logger.Error("failed to provision user from IdP", "provider", providerID, "error", err)
|
||||
http.Redirect(w, r, "/login?error=provisioning_failed", http.StatusTemporaryRedirect)
|
||||
return
|
||||
}
|
||||
|
||||
// Create session
|
||||
session, err := h.sessionStore.CreateSession(r.Context(), user.ID, 24*time.Hour)
|
||||
if err != nil {
|
||||
h.logger.Error("failed to create session after IdP login", "error", err)
|
||||
http.Error(w, "Failed to create session", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
// Set session cookie
|
||||
http.SetCookie(w, &http.Cookie{
|
||||
Name: auth.SessionCookieName,
|
||||
Value: session.SessionID,
|
||||
Path: "/",
|
||||
HttpOnly: true,
|
||||
SameSite: http.SameSiteLaxMode,
|
||||
MaxAge: 86400, // 24 hours
|
||||
})
|
||||
|
||||
auth.LogAuthEvent(r.Context(), h.logger, auth.AuthEvent{
|
||||
Type: auth.EventLoginSuccess,
|
||||
UserID: user.ID,
|
||||
Username: user.Username,
|
||||
RemoteIP: r.RemoteAddr,
|
||||
Details: map[string]any{
|
||||
"provider": providerID,
|
||||
"external_id": extUser.ExternalID,
|
||||
},
|
||||
})
|
||||
|
||||
// Ensure human agent and channel in background
|
||||
if h.agentProvisioner != nil {
|
||||
go func() {
|
||||
ctx := context.Background()
|
||||
if err := h.agentProvisioner.ProvisionHumanAgent(ctx, user.Username, user.DisplayName, user.ID); err != nil {
|
||||
h.logger.Warn("failed to ensure human agent after IdP login", "username", user.Username, "error", err)
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
// Redirect to home
|
||||
http.Redirect(w, r, "/", http.StatusTemporaryRedirect)
|
||||
}
|
||||
|
||||
// findOrCreateUser looks up or provisions a local user from external identity.
|
||||
// 1. Check user_identities for (provider, external_id) -> if found, load user
|
||||
// 2. If not found, check users by email -> if found, link identity
|
||||
// 3. If no user at all, create new user with random password
|
||||
func (h *Handlers) findOrCreateUser(ctx context.Context, ext *ExternalUser) (*auth.User, error) {
|
||||
// Step 1: Check existing identity link
|
||||
userID, err := h.idStore.FindByProvider(ctx, ext.ProviderID, ext.ExternalID)
|
||||
if err == nil {
|
||||
// Found existing link
|
||||
user, err := h.userStore.GetUserByID(ctx, userID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("load linked user %d: %w", userID, err)
|
||||
}
|
||||
return user, nil
|
||||
}
|
||||
if err != sql.ErrNoRows {
|
||||
return nil, fmt.Errorf("find identity: %w", err)
|
||||
}
|
||||
|
||||
// Step 2: Try to find user by email
|
||||
if ext.Email != "" {
|
||||
user, err := h.userStore.GetUserByEmail(ctx, ext.Email)
|
||||
if err == nil {
|
||||
// Found user with matching email — link this identity
|
||||
if linkErr := h.idStore.Create(ctx, user.ID, ext.ProviderID, ext.ExternalID, ext.Email, ext.Name, ext.RawClaims); linkErr != nil {
|
||||
h.logger.Warn("failed to link identity to existing user", "user_id", user.ID, "error", linkErr)
|
||||
}
|
||||
return user, nil
|
||||
}
|
||||
// If error is not "not found", it's a real error
|
||||
if err != auth.ErrUserNotFound {
|
||||
return nil, fmt.Errorf("find user by email: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Step 3: Create new user
|
||||
username := sanitizeUsername(ext.Username, ext.ProviderID, ext.ExternalID)
|
||||
displayName := ext.Name
|
||||
if displayName == "" {
|
||||
displayName = username
|
||||
}
|
||||
|
||||
// Generate random password (user will login via IdP)
|
||||
randomPW, err := generateRandomPassword()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("generate password: %w", err)
|
||||
}
|
||||
|
||||
user, err := h.userStore.CreateUser(ctx, username, randomPW, displayName)
|
||||
if err != nil {
|
||||
// If username conflicts, try with a suffix
|
||||
if strings.Contains(err.Error(), "already exists") {
|
||||
suffix, _ := generateShortID()
|
||||
username = username + "_" + suffix
|
||||
user, err = h.userStore.CreateUser(ctx, username, randomPW, displayName)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create user: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Set email on the user if available
|
||||
if ext.Email != "" {
|
||||
if emailStore, ok := h.userStore.(*auth.SQLiteUserStore); ok {
|
||||
emailStore.SetEmail(ctx, user.ID, ext.Email)
|
||||
}
|
||||
}
|
||||
|
||||
// Link identity
|
||||
if linkErr := h.idStore.Create(ctx, user.ID, ext.ProviderID, ext.ExternalID, ext.Email, ext.Name, ext.RawClaims); linkErr != nil {
|
||||
h.logger.Warn("failed to link identity to new user", "user_id", user.ID, "error", linkErr)
|
||||
}
|
||||
|
||||
h.logger.Info("provisioned new user from IdP",
|
||||
"username", username,
|
||||
"provider", ext.ProviderID,
|
||||
"external_id", ext.ExternalID,
|
||||
"email", ext.Email,
|
||||
)
|
||||
|
||||
return user, nil
|
||||
}
|
||||
|
||||
// sanitizeUsername normalizes a username from an IdP to match SynapBus requirements.
|
||||
// Usernames must be 3-64 chars, alphanumeric and underscore only.
|
||||
func sanitizeUsername(username, provider, externalID string) string {
|
||||
if username == "" {
|
||||
username = provider + "_" + externalID
|
||||
}
|
||||
|
||||
// Replace invalid characters with underscore
|
||||
var result strings.Builder
|
||||
for _, r := range username {
|
||||
if (r >= 'a' && r <= 'z') || (r >= 'A' && r <= 'Z') || (r >= '0' && r <= '9') || r == '_' {
|
||||
result.WriteRune(r)
|
||||
} else {
|
||||
result.WriteRune('_')
|
||||
}
|
||||
}
|
||||
username = result.String()
|
||||
|
||||
// Trim underscores from edges and ensure minimum length
|
||||
username = strings.Trim(username, "_")
|
||||
if len(username) < 3 {
|
||||
username = username + "_user"
|
||||
}
|
||||
if len(username) > 64 {
|
||||
username = username[:64]
|
||||
}
|
||||
|
||||
return username
|
||||
}
|
||||
|
||||
// generateState creates a cryptographically random state parameter for OAuth.
|
||||
func generateState() (string, error) {
|
||||
b := make([]byte, 16)
|
||||
if _, err := rand.Read(b); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return hex.EncodeToString(b), nil
|
||||
}
|
||||
|
||||
// generateRandomPassword creates a random password for IdP-provisioned users.
|
||||
func generateRandomPassword() (string, error) {
|
||||
b := make([]byte, 32)
|
||||
if _, err := rand.Read(b); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return hex.EncodeToString(b), nil
|
||||
}
|
||||
|
||||
// generateShortID creates a short random string for username disambiguation.
|
||||
func generateShortID() (string, error) {
|
||||
b := make([]byte, 3)
|
||||
if _, err := rand.Read(b); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return hex.EncodeToString(b), nil
|
||||
}
|
||||
@@ -0,0 +1,392 @@
|
||||
package idp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
|
||||
"github.com/synapbus/synapbus/internal/auth"
|
||||
"github.com/synapbus/synapbus/internal/storage"
|
||||
)
|
||||
|
||||
func newTestDB(t *testing.T) *sql.DB {
|
||||
t.Helper()
|
||||
dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", t.Name())
|
||||
db, err := sql.Open("sqlite", dsn)
|
||||
if err != nil {
|
||||
t.Fatalf("open database: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { db.Close() })
|
||||
|
||||
if _, err := db.Exec("PRAGMA foreign_keys=ON"); err != nil {
|
||||
t.Fatalf("enable foreign keys: %v", err)
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
if err := storage.RunMigrations(ctx, db); err != nil {
|
||||
t.Fatalf("run migrations: %v", err)
|
||||
}
|
||||
|
||||
return db
|
||||
}
|
||||
|
||||
func newTestUserStore(db *sql.DB, t *testing.T) auth.UserStore {
|
||||
t.Helper()
|
||||
return auth.NewSQLiteUserStore(db, 10) // low cost for fast tests
|
||||
}
|
||||
|
||||
// --- GitHub provider tests ---
|
||||
|
||||
func TestGitHubProvider_AuthCodeURL(t *testing.T) {
|
||||
p := NewGitHubProvider("test-client-id", "test-secret", "http://localhost:8080/auth/callback/github")
|
||||
|
||||
url := p.AuthCodeURL("test-state-123")
|
||||
|
||||
if url == "" {
|
||||
t.Fatal("AuthCodeURL returned empty string")
|
||||
}
|
||||
if p.ID() != "github" {
|
||||
t.Errorf("ID() = %q, want %q", p.ID(), "github")
|
||||
}
|
||||
if p.Type() != "oauth" {
|
||||
t.Errorf("Type() = %q, want %q", p.Type(), "oauth")
|
||||
}
|
||||
if p.DisplayName() != "GitHub" {
|
||||
t.Errorf("DisplayName() = %q, want %q", p.DisplayName(), "GitHub")
|
||||
}
|
||||
|
||||
// URL should contain the client ID and state
|
||||
if got := url; got == "" {
|
||||
t.Error("expected non-empty URL")
|
||||
}
|
||||
}
|
||||
|
||||
// --- OIDC domain validation tests ---
|
||||
|
||||
func TestOIDCProvider_DomainRestriction(t *testing.T) {
|
||||
p := &OIDCProvider{
|
||||
id: "test",
|
||||
displayName: "Test",
|
||||
allowedDomains: []string{"gcore.com", "example.com"},
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
email string
|
||||
hd string
|
||||
wantError bool
|
||||
}{
|
||||
{"allowed domain via hd", "user@gcore.com", "gcore.com", false},
|
||||
{"allowed domain via email", "user@example.com", "", false},
|
||||
{"disallowed domain", "user@evil.com", "evil.com", true},
|
||||
{"no domain info", "", "", true},
|
||||
{"case insensitive", "User@Gcore.COM", "Gcore.COM", false},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
err := p.validateDomain(tt.email, tt.hd)
|
||||
if (err != nil) != tt.wantError {
|
||||
t.Errorf("validateDomain(%q, %q) error = %v, wantError %v", tt.email, tt.hd, err, tt.wantError)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestOIDCProvider_NoDomainRestriction(t *testing.T) {
|
||||
p := &OIDCProvider{
|
||||
id: "test",
|
||||
displayName: "Test",
|
||||
allowedDomains: nil,
|
||||
}
|
||||
|
||||
if len(p.allowedDomains) != 0 {
|
||||
t.Error("expected empty allowed domains")
|
||||
}
|
||||
}
|
||||
|
||||
// --- Store tests ---
|
||||
|
||||
func TestUserIdentityStore_CreateAndFind(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewUserIdentityStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
// Create a user first
|
||||
_, err := db.ExecContext(ctx,
|
||||
`INSERT INTO users (username, password_hash, display_name, role) VALUES (?, ?, ?, ?)`,
|
||||
"testuser", "hash", "Test User", "user",
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("create user: %v", err)
|
||||
}
|
||||
|
||||
var userID int64
|
||||
db.QueryRowContext(ctx, "SELECT id FROM users WHERE username = ?", "testuser").Scan(&userID)
|
||||
|
||||
// Create identity
|
||||
claims := map[string]any{"sub": "12345", "login": "ghuser"}
|
||||
err = store.Create(ctx, userID, "github", "12345", "test@example.com", "Test User", claims)
|
||||
if err != nil {
|
||||
t.Fatalf("Create: %v", err)
|
||||
}
|
||||
|
||||
// Find by provider
|
||||
foundID, err := store.FindByProvider(ctx, "github", "12345")
|
||||
if err != nil {
|
||||
t.Fatalf("FindByProvider: %v", err)
|
||||
}
|
||||
if foundID != userID {
|
||||
t.Errorf("FindByProvider = %d, want %d", foundID, userID)
|
||||
}
|
||||
|
||||
// Find non-existent
|
||||
_, err = store.FindByProvider(ctx, "github", "99999")
|
||||
if err != sql.ErrNoRows {
|
||||
t.Errorf("expected sql.ErrNoRows, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUserIdentityStore_ListByUser(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewUserIdentityStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
_, err := db.ExecContext(ctx,
|
||||
`INSERT INTO users (username, password_hash, display_name, role) VALUES (?, ?, ?, ?)`,
|
||||
"multiuser", "hash", "Multi User", "user",
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("create user: %v", err)
|
||||
}
|
||||
|
||||
var userID int64
|
||||
db.QueryRowContext(ctx, "SELECT id FROM users WHERE username = ?", "multiuser").Scan(&userID)
|
||||
|
||||
store.Create(ctx, userID, "github", "gh-123", "user@gh.com", "GH User", map[string]any{})
|
||||
store.Create(ctx, userID, "google", "goog-456", "user@google.com", "Google User", map[string]any{})
|
||||
|
||||
identities, err := store.ListByUser(ctx, userID)
|
||||
if err != nil {
|
||||
t.Fatalf("ListByUser: %v", err)
|
||||
}
|
||||
if len(identities) != 2 {
|
||||
t.Errorf("got %d identities, want 2", len(identities))
|
||||
}
|
||||
|
||||
if identities[0].Provider != "github" {
|
||||
t.Errorf("first identity provider = %q, want %q", identities[0].Provider, "github")
|
||||
}
|
||||
if identities[1].Provider != "google" {
|
||||
t.Errorf("second identity provider = %q, want %q", identities[1].Provider, "google")
|
||||
}
|
||||
|
||||
empty, err := store.ListByUser(ctx, 99999)
|
||||
if err != nil {
|
||||
t.Fatalf("ListByUser (empty): %v", err)
|
||||
}
|
||||
if len(empty) != 0 {
|
||||
t.Errorf("expected empty list, got %d", len(empty))
|
||||
}
|
||||
}
|
||||
|
||||
// --- Handler tests ---
|
||||
|
||||
func TestHandleListProviders(t *testing.T) {
|
||||
providers := []Provider{
|
||||
NewGitHubProvider("id", "secret", "http://localhost/callback"),
|
||||
&mockProvider{id: "google", providerType: "oidc", name: "Google"},
|
||||
}
|
||||
|
||||
handlers := NewHandlers(providers, nil, nil, nil, nil)
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/auth/providers", nil)
|
||||
w := httptest.NewRecorder()
|
||||
|
||||
handlers.HandleListProviders(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want %d", w.Code, http.StatusOK)
|
||||
}
|
||||
|
||||
var resp struct {
|
||||
Providers []ProviderInfo `json:"providers"`
|
||||
}
|
||||
if err := json.NewDecoder(w.Body).Decode(&resp); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
|
||||
if len(resp.Providers) != 2 {
|
||||
t.Fatalf("got %d providers, want 2", len(resp.Providers))
|
||||
}
|
||||
|
||||
if resp.Providers[0].ID != "github" {
|
||||
t.Errorf("first provider ID = %q, want %q", resp.Providers[0].ID, "github")
|
||||
}
|
||||
if resp.Providers[0].DisplayName != "GitHub" {
|
||||
t.Errorf("first provider DisplayName = %q, want %q", resp.Providers[0].DisplayName, "GitHub")
|
||||
}
|
||||
if resp.Providers[1].ID != "google" {
|
||||
t.Errorf("second provider ID = %q, want %q", resp.Providers[1].ID, "google")
|
||||
}
|
||||
}
|
||||
|
||||
// --- User provisioning tests ---
|
||||
|
||||
func TestSanitizeUsername(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
username string
|
||||
provider string
|
||||
externalID string
|
||||
wantMin string // exact match, or empty for length-only check
|
||||
maxLen int
|
||||
}{
|
||||
{"normal", "testuser", "github", "123", "testuser", 0},
|
||||
{"with dash", "test-user", "github", "123", "test_user", 0},
|
||||
{"with dots", "test.user", "github", "123", "test_user", 0},
|
||||
{"with at sign", "user@email.com", "github", "123", "user_email_com", 0},
|
||||
{"too short", "ab", "github", "123", "ab_user", 0},
|
||||
{"empty", "", "github", "123", "github_123", 0},
|
||||
{"too long", "", "", "", "", 64},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
input := tt.username
|
||||
if tt.name == "too long" {
|
||||
b := make([]byte, 100)
|
||||
for i := range b {
|
||||
b[i] = 'a'
|
||||
}
|
||||
input = string(b)
|
||||
}
|
||||
got := sanitizeUsername(input, tt.provider, tt.externalID)
|
||||
if tt.maxLen > 0 {
|
||||
if len(got) > tt.maxLen {
|
||||
t.Errorf("sanitizeUsername too long: len=%d, want <= %d", len(got), tt.maxLen)
|
||||
}
|
||||
return
|
||||
}
|
||||
if got != tt.wantMin {
|
||||
t.Errorf("sanitizeUsername(%q, %q, %q) = %q, want %q", tt.username, tt.provider, tt.externalID, got, tt.wantMin)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestFindOrCreateUser_NewUser(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
idStore := NewUserIdentityStore(db)
|
||||
userStore := newTestUserStore(db, t)
|
||||
|
||||
handlers := &Handlers{
|
||||
idStore: idStore,
|
||||
userStore: userStore,
|
||||
logger: slog.Default(),
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
ext := &ExternalUser{
|
||||
ProviderID: "github",
|
||||
ExternalID: "gh-newuser-001",
|
||||
Email: "newuser@example.com",
|
||||
Name: "New User",
|
||||
Username: "newghuser",
|
||||
RawClaims: map[string]any{"login": "newghuser"},
|
||||
}
|
||||
|
||||
user, err := handlers.findOrCreateUser(ctx, ext)
|
||||
if err != nil {
|
||||
t.Fatalf("findOrCreateUser: %v", err)
|
||||
}
|
||||
|
||||
if user.Username != "newghuser" {
|
||||
t.Errorf("Username = %q, want %q", user.Username, "newghuser")
|
||||
}
|
||||
if user.DisplayName != "New User" {
|
||||
t.Errorf("DisplayName = %q, want %q", user.DisplayName, "New User")
|
||||
}
|
||||
|
||||
// Identity should be linked
|
||||
foundID, err := idStore.FindByProvider(ctx, "github", "gh-newuser-001")
|
||||
if err != nil {
|
||||
t.Fatalf("identity not linked: %v", err)
|
||||
}
|
||||
if foundID != user.ID {
|
||||
t.Errorf("linked user ID = %d, want %d", foundID, user.ID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFindOrCreateUser_ExistingIdentity(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
idStore := NewUserIdentityStore(db)
|
||||
userStore := newTestUserStore(db, t)
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Create a user first
|
||||
user, err := userStore.CreateUser(ctx, "existinguser", "password1234", "Existing User")
|
||||
if err != nil {
|
||||
t.Fatalf("CreateUser: %v", err)
|
||||
}
|
||||
|
||||
// Link identity manually
|
||||
err = idStore.Create(ctx, user.ID, "github", "gh-existing-001", "existing@example.com", "Existing", map[string]any{})
|
||||
if err != nil {
|
||||
t.Fatalf("Create identity: %v", err)
|
||||
}
|
||||
|
||||
handlers := &Handlers{
|
||||
idStore: idStore,
|
||||
userStore: userStore,
|
||||
logger: slog.Default(),
|
||||
}
|
||||
|
||||
ext := &ExternalUser{
|
||||
ProviderID: "github",
|
||||
ExternalID: "gh-existing-001",
|
||||
Email: "existing@example.com",
|
||||
Name: "Existing User",
|
||||
Username: "existinguser",
|
||||
}
|
||||
|
||||
found, err := handlers.findOrCreateUser(ctx, ext)
|
||||
if err != nil {
|
||||
t.Fatalf("findOrCreateUser: %v", err)
|
||||
}
|
||||
if found.ID != user.ID {
|
||||
t.Errorf("found user ID = %d, want %d", found.ID, user.ID)
|
||||
}
|
||||
}
|
||||
|
||||
// --- Mock helpers ---
|
||||
|
||||
type mockProvider struct {
|
||||
id string
|
||||
providerType string
|
||||
name string
|
||||
}
|
||||
|
||||
func (m *mockProvider) ID() string { return m.id }
|
||||
func (m *mockProvider) Type() string { return m.providerType }
|
||||
func (m *mockProvider) DisplayName() string { return m.name }
|
||||
func (m *mockProvider) AuthCodeURL(state string) string {
|
||||
return "https://mock.idp.example.com/authorize?state=" + state
|
||||
}
|
||||
func (m *mockProvider) Exchange(ctx context.Context, code string) (*ExternalUser, error) {
|
||||
return &ExternalUser{
|
||||
ProviderID: m.id,
|
||||
ExternalID: "mock-ext-id",
|
||||
Email: "mock@example.com",
|
||||
Name: "Mock User",
|
||||
Username: "mockuser",
|
||||
}, nil
|
||||
}
|
||||
@@ -0,0 +1,174 @@
|
||||
package idp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/coreos/go-oidc/v3/oidc"
|
||||
"golang.org/x/oauth2"
|
||||
)
|
||||
|
||||
// OIDCProvider implements OpenID Connect authentication.
|
||||
// Works with any OIDC-compliant provider (Google, Azure AD, etc.).
|
||||
type OIDCProvider struct {
|
||||
id string
|
||||
displayName string
|
||||
issuerURL string
|
||||
config *oauth2.Config
|
||||
verifier *oidc.IDTokenVerifier
|
||||
allowedDomains []string // Empty means allow all domains
|
||||
groupMapping map[string]string
|
||||
}
|
||||
|
||||
// OIDCConfig holds configuration for an OIDC provider.
|
||||
type OIDCConfig struct {
|
||||
ID string
|
||||
DisplayName string
|
||||
IssuerURL string
|
||||
ClientID string
|
||||
ClientSecret string
|
||||
RedirectURL string
|
||||
Scopes []string
|
||||
AllowedDomains []string
|
||||
GroupMapping map[string]string
|
||||
}
|
||||
|
||||
// NewOIDCProvider creates a new OIDC provider using discovery.
|
||||
func NewOIDCProvider(ctx context.Context, cfg OIDCConfig) (*OIDCProvider, error) {
|
||||
provider, err := oidc.NewProvider(ctx, cfg.IssuerURL)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("oidc: discover %s: %w", cfg.IssuerURL, err)
|
||||
}
|
||||
|
||||
scopes := cfg.Scopes
|
||||
if len(scopes) == 0 {
|
||||
scopes = []string{oidc.ScopeOpenID, "email", "profile"}
|
||||
}
|
||||
|
||||
oauthConfig := &oauth2.Config{
|
||||
ClientID: cfg.ClientID,
|
||||
ClientSecret: cfg.ClientSecret,
|
||||
RedirectURL: cfg.RedirectURL,
|
||||
Endpoint: provider.Endpoint(),
|
||||
Scopes: scopes,
|
||||
}
|
||||
|
||||
verifier := provider.Verifier(&oidc.Config{
|
||||
ClientID: cfg.ClientID,
|
||||
})
|
||||
|
||||
return &OIDCProvider{
|
||||
id: cfg.ID,
|
||||
displayName: cfg.DisplayName,
|
||||
issuerURL: cfg.IssuerURL,
|
||||
config: oauthConfig,
|
||||
verifier: verifier,
|
||||
allowedDomains: cfg.AllowedDomains,
|
||||
groupMapping: cfg.GroupMapping,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (p *OIDCProvider) ID() string { return p.id }
|
||||
func (p *OIDCProvider) Type() string { return "oidc" }
|
||||
func (p *OIDCProvider) DisplayName() string { return p.displayName }
|
||||
|
||||
func (p *OIDCProvider) AuthCodeURL(state string) string {
|
||||
return p.config.AuthCodeURL(state)
|
||||
}
|
||||
|
||||
func (p *OIDCProvider) Exchange(ctx context.Context, code string) (*ExternalUser, error) {
|
||||
token, err := p.config.Exchange(ctx, code)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("oidc: exchange code: %w", err)
|
||||
}
|
||||
|
||||
rawIDToken, ok := token.Extra("id_token").(string)
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("oidc: no id_token in response")
|
||||
}
|
||||
|
||||
idToken, err := p.verifier.Verify(ctx, rawIDToken)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("oidc: verify id_token: %w", err)
|
||||
}
|
||||
|
||||
// Extract claims
|
||||
var claims struct {
|
||||
Sub string `json:"sub"`
|
||||
Email string `json:"email"`
|
||||
Name string `json:"name"`
|
||||
Username string `json:"preferred_username"`
|
||||
HD string `json:"hd"` // Google hosted domain
|
||||
Groups []string `json:"groups"`
|
||||
}
|
||||
if err := idToken.Claims(&claims); err != nil {
|
||||
return nil, fmt.Errorf("oidc: parse claims: %w", err)
|
||||
}
|
||||
|
||||
// Extract raw claims for storage
|
||||
var rawClaims map[string]any
|
||||
if err := idToken.Claims(&rawClaims); err != nil {
|
||||
rawClaims = map[string]any{"sub": claims.Sub}
|
||||
}
|
||||
|
||||
// Validate domain restriction
|
||||
if len(p.allowedDomains) > 0 {
|
||||
if err := p.validateDomain(claims.Email, claims.HD); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
// Map groups through group mapping
|
||||
var mappedGroups []string
|
||||
if len(p.groupMapping) > 0 {
|
||||
for _, group := range claims.Groups {
|
||||
if mapped, ok := p.groupMapping[group]; ok {
|
||||
mappedGroups = append(mappedGroups, mapped)
|
||||
}
|
||||
}
|
||||
} else {
|
||||
mappedGroups = claims.Groups
|
||||
}
|
||||
|
||||
// Derive username from email if preferred_username is empty
|
||||
username := claims.Username
|
||||
if username == "" && claims.Email != "" {
|
||||
parts := strings.SplitN(claims.Email, "@", 2)
|
||||
username = parts[0]
|
||||
}
|
||||
|
||||
return &ExternalUser{
|
||||
ProviderID: p.id,
|
||||
ExternalID: claims.Sub,
|
||||
Email: claims.Email,
|
||||
Name: claims.Name,
|
||||
Username: username,
|
||||
Groups: mappedGroups,
|
||||
RawClaims: rawClaims,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// validateDomain checks that the user's email domain is in the allowed list.
|
||||
func (p *OIDCProvider) validateDomain(email, hostedDomain string) error {
|
||||
// Prefer hd (hosted domain) claim when available (Google Workspace)
|
||||
domain := hostedDomain
|
||||
if domain == "" && email != "" {
|
||||
parts := strings.SplitN(email, "@", 2)
|
||||
if len(parts) == 2 {
|
||||
domain = parts[1]
|
||||
}
|
||||
}
|
||||
|
||||
if domain == "" {
|
||||
return fmt.Errorf("oidc: no domain found in claims, required domains: %v", p.allowedDomains)
|
||||
}
|
||||
|
||||
for _, allowed := range p.allowedDomains {
|
||||
if strings.EqualFold(domain, allowed) {
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
return fmt.Errorf("oidc: domain %q not in allowed list %v", domain, p.allowedDomains)
|
||||
}
|
||||
@@ -0,0 +1,37 @@
|
||||
// Package idp provides enterprise identity provider integrations for SynapBus.
|
||||
// Supported providers: GitHub (OAuth), Google (OIDC), Azure AD (OIDC).
|
||||
package idp
|
||||
|
||||
import "context"
|
||||
|
||||
// Provider is the interface that all identity providers must implement.
|
||||
type Provider interface {
|
||||
// ID returns the unique identifier for this provider (e.g., "github", "google", "azuread").
|
||||
ID() string
|
||||
// Type returns the provider type ("oauth" or "oidc").
|
||||
Type() string
|
||||
// DisplayName returns the human-readable name (e.g., "GitHub", "Google").
|
||||
DisplayName() string
|
||||
// AuthCodeURL generates the authorization URL with the given state parameter.
|
||||
AuthCodeURL(state string) string
|
||||
// Exchange trades an authorization code for user information.
|
||||
Exchange(ctx context.Context, code string) (*ExternalUser, error)
|
||||
}
|
||||
|
||||
// ExternalUser represents user information obtained from an external identity provider.
|
||||
type ExternalUser struct {
|
||||
ProviderID string
|
||||
ExternalID string
|
||||
Email string
|
||||
Name string
|
||||
Username string
|
||||
Groups []string
|
||||
RawClaims map[string]any
|
||||
}
|
||||
|
||||
// ProviderInfo is a minimal representation of a provider for API responses.
|
||||
type ProviderInfo struct {
|
||||
ID string `json:"id"`
|
||||
Type string `json:"type"`
|
||||
DisplayName string `json:"display_name"`
|
||||
}
|
||||
@@ -0,0 +1,97 @@
|
||||
package idp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"time"
|
||||
)
|
||||
|
||||
// UserIdentity represents an external identity linked to a local user.
|
||||
type UserIdentity struct {
|
||||
ID int64 `json:"id"`
|
||||
UserID int64 `json:"user_id"`
|
||||
Provider string `json:"provider"`
|
||||
ExternalID string `json:"external_id"`
|
||||
Email string `json:"email,omitempty"`
|
||||
DisplayName string `json:"display_name,omitempty"`
|
||||
RawClaims string `json:"raw_claims"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
// UserIdentityStore manages the user_identities table.
|
||||
type UserIdentityStore struct {
|
||||
db *sql.DB
|
||||
}
|
||||
|
||||
// NewUserIdentityStore creates a new identity store backed by SQLite.
|
||||
func NewUserIdentityStore(db *sql.DB) *UserIdentityStore {
|
||||
return &UserIdentityStore{db: db}
|
||||
}
|
||||
|
||||
// FindByProvider looks up a user ID by provider name and external ID.
|
||||
// Returns sql.ErrNoRows if not found.
|
||||
func (s *UserIdentityStore) FindByProvider(ctx context.Context, provider, externalID string) (int64, error) {
|
||||
var userID int64
|
||||
err := s.db.QueryRowContext(ctx,
|
||||
`SELECT user_id FROM user_identities WHERE provider = ? AND external_id = ?`,
|
||||
provider, externalID,
|
||||
).Scan(&userID)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return userID, nil
|
||||
}
|
||||
|
||||
// Create links an external identity to a local user.
|
||||
func (s *UserIdentityStore) Create(ctx context.Context, userID int64, provider, externalID, email, displayName string, rawClaims map[string]any) error {
|
||||
claimsJSON, err := json.Marshal(rawClaims)
|
||||
if err != nil {
|
||||
claimsJSON = []byte("{}")
|
||||
}
|
||||
|
||||
_, err = s.db.ExecContext(ctx,
|
||||
`INSERT INTO user_identities (user_id, provider, external_id, email, display_name, raw_claims, created_at, updated_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`,
|
||||
userID, provider, externalID, email, displayName, string(claimsJSON),
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("create identity: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ListByUser returns all external identities linked to a user.
|
||||
func (s *UserIdentityStore) ListByUser(ctx context.Context, userID int64) ([]UserIdentity, error) {
|
||||
rows, err := s.db.QueryContext(ctx,
|
||||
`SELECT id, user_id, provider, external_id, email, display_name, raw_claims, created_at, updated_at
|
||||
FROM user_identities WHERE user_id = ? ORDER BY created_at`,
|
||||
userID,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list identities: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var identities []UserIdentity
|
||||
for rows.Next() {
|
||||
var identity UserIdentity
|
||||
var email, displayName sql.NullString
|
||||
if err := rows.Scan(
|
||||
&identity.ID, &identity.UserID, &identity.Provider,
|
||||
&identity.ExternalID, &email, &displayName,
|
||||
&identity.RawClaims, &identity.CreatedAt, &identity.UpdatedAt,
|
||||
); err != nil {
|
||||
return nil, fmt.Errorf("scan identity: %w", err)
|
||||
}
|
||||
identity.Email = email.String
|
||||
identity.DisplayName = displayName.String
|
||||
identities = append(identities, identity)
|
||||
}
|
||||
if identities == nil {
|
||||
identities = []UserIdentity{}
|
||||
}
|
||||
return identities, rows.Err()
|
||||
}
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/ory/fosite"
|
||||
)
|
||||
@@ -107,8 +108,11 @@ func RequireBearer(provider fosite.OAuth2Provider, userStore UserStore) func(htt
|
||||
token := parts[1]
|
||||
_ = token
|
||||
|
||||
// Use fosite introspection
|
||||
_, ar, err := provider.IntrospectToken(r.Context(), parts[1], fosite.AccessToken, new(fositeSession))
|
||||
// Use fosite introspection — decouple from HTTP request context so
|
||||
// token validation completes even if the client disconnects.
|
||||
dbCtx, cancel := context.WithTimeout(context.WithoutCancel(r.Context()), 10*time.Second)
|
||||
_, ar, err := provider.IntrospectToken(dbCtx, parts[1], fosite.AccessToken, new(fositeSession))
|
||||
cancel()
|
||||
if err != nil {
|
||||
slog.Debug("bearer token validation failed", "error", err)
|
||||
w.Header().Set("WWW-Authenticate", `Bearer error="invalid_token"`)
|
||||
|
||||
@@ -18,10 +18,12 @@ type UserStore interface {
|
||||
CreateUser(ctx context.Context, username, password, displayName string) (*User, error)
|
||||
GetUserByID(ctx context.Context, id int64) (*User, error)
|
||||
GetUserByUsername(ctx context.Context, username string) (*User, error)
|
||||
GetUserByEmail(ctx context.Context, email string) (*User, error)
|
||||
UpdatePassword(ctx context.Context, userID int64, newPassword string) error
|
||||
ListUsers(ctx context.Context) ([]*User, error)
|
||||
CountUsers(ctx context.Context) (int, error)
|
||||
VerifyPassword(ctx context.Context, username, password string) (*User, error)
|
||||
UpdateDisplayName(ctx context.Context, userID int64, displayName string) error
|
||||
}
|
||||
|
||||
// SQLiteUserStore implements UserStore using SQLite.
|
||||
@@ -115,6 +117,36 @@ func (s *SQLiteUserStore) GetUserByID(ctx context.Context, id int64) (*User, err
|
||||
return user, nil
|
||||
}
|
||||
|
||||
// GetUserByEmail retrieves a user by their email address.
|
||||
// Returns ErrUserNotFound if no user has the given email or if email is empty.
|
||||
func (s *SQLiteUserStore) GetUserByEmail(ctx context.Context, email string) (*User, error) {
|
||||
if email == "" {
|
||||
return nil, ErrUserNotFound
|
||||
}
|
||||
user := &User{}
|
||||
err := s.db.QueryRowContext(ctx,
|
||||
`SELECT id, username, password_hash, display_name, role, created_at, updated_at
|
||||
FROM users WHERE email = ?`, email,
|
||||
).Scan(&user.ID, &user.Username, &user.PasswordHash, &user.DisplayName,
|
||||
&user.Role, &user.CreatedAt, &user.UpdatedAt)
|
||||
if err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, ErrUserNotFound
|
||||
}
|
||||
return nil, fmt.Errorf("query user by email: %w", err)
|
||||
}
|
||||
return user, nil
|
||||
}
|
||||
|
||||
// SetEmail updates a user's email address.
|
||||
func (s *SQLiteUserStore) SetEmail(ctx context.Context, userID int64, email string) error {
|
||||
_, err := s.db.ExecContext(ctx,
|
||||
`UPDATE users SET email = ?, updated_at = CURRENT_TIMESTAMP WHERE id = ?`,
|
||||
email, userID,
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
// GetUserByUsername retrieves a user by their username.
|
||||
func (s *SQLiteUserStore) GetUserByUsername(ctx context.Context, username string) (*User, error) {
|
||||
user := &User{}
|
||||
@@ -215,3 +247,26 @@ func (s *SQLiteUserStore) VerifyPassword(ctx context.Context, username, password
|
||||
|
||||
return user, nil
|
||||
}
|
||||
|
||||
// UpdateDisplayName changes a user's display name.
|
||||
func (s *SQLiteUserStore) UpdateDisplayName(ctx context.Context, userID int64, displayName string) error {
|
||||
displayName = strings.TrimSpace(displayName)
|
||||
if displayName == "" {
|
||||
return fmt.Errorf("display name cannot be empty")
|
||||
}
|
||||
result, err := s.db.ExecContext(ctx,
|
||||
`UPDATE users SET display_name = ?, updated_at = CURRENT_TIMESTAMP WHERE id = ?`,
|
||||
displayName, userID,
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("update display name: %w", err)
|
||||
}
|
||||
rows, err := result.RowsAffected()
|
||||
if err != nil {
|
||||
return fmt.Errorf("check rows affected: %w", err)
|
||||
}
|
||||
if rows == 0 {
|
||||
return ErrUserNotFound
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -459,25 +459,55 @@ func (s *Service) UpdateChannel(ctx context.Context, channelID int64, req Update
|
||||
return ch, nil
|
||||
}
|
||||
|
||||
// UpdateChannelSettings updates the workflow-related settings for a channel.
|
||||
func (s *Service) UpdateChannelSettings(ctx context.Context, channelID int64, settings ChannelSettings) (*Channel, error) {
|
||||
store, ok := s.store.(*SQLiteChannelStore)
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("channel store does not support settings update")
|
||||
}
|
||||
|
||||
if err := store.UpdateChannelSettings(ctx, channelID, settings); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Reload channel to return updated state
|
||||
ch, err := s.store.GetChannel(ctx, channelID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
s.logger.Info("channel settings updated",
|
||||
"channel_id", channelID,
|
||||
"auto_approve", settings.AutoApprove,
|
||||
)
|
||||
|
||||
return ch, nil
|
||||
}
|
||||
|
||||
// BroadcastMessage sends a message to a channel. It creates a single channel
|
||||
// message (visible in the channel timeline via GetChannelMessages) and also
|
||||
// delivers individual DM notifications to each member's inbox.
|
||||
// If the message body contains @mentions, mentioned members receive a
|
||||
// "mention":true flag in their inbox notification metadata, and the channel
|
||||
// message metadata includes "mentioned_agents".
|
||||
func (s *Service) BroadcastMessage(ctx context.Context, channelID int64, fromAgent, body string, priority int, metadata string) ([]*messaging.Message, error) {
|
||||
func (s *Service) BroadcastMessage(ctx context.Context, channelID int64, fromAgent, body string, priority int, metadata string, replyTo *int64, attachments []string) ([]*messaging.Message, error) {
|
||||
ch, err := s.store.GetChannel(ctx, channelID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Verify sender is a member
|
||||
// Verify sender is a member; auto-join public channels on first send.
|
||||
isMember, err := s.store.IsMember(ctx, channelID, fromAgent)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("check membership: %w", err)
|
||||
}
|
||||
if !isMember {
|
||||
return nil, ErrNotChannelMember
|
||||
if ch.IsPrivate {
|
||||
return nil, ErrNotChannelMember
|
||||
}
|
||||
if err := s.JoinChannel(ctx, channelID, fromAgent); err != nil {
|
||||
return nil, fmt.Errorf("auto-join public channel: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Get members for mentions and inbox notifications
|
||||
@@ -517,19 +547,21 @@ func (s *Service) BroadcastMessage(ctx context.Context, channelID int64, fromAge
|
||||
channelMetaBytes, _ := json.Marshal(channelMetaObj)
|
||||
|
||||
channelMsg, err := s.msgService.SendMessage(ctx, fromAgent, "", body, messaging.SendOptions{
|
||||
Subject: fmt.Sprintf("channel:%s", ch.Name),
|
||||
Priority: priority,
|
||||
Metadata: string(channelMetaBytes),
|
||||
ChannelID: &channelID,
|
||||
Subject: fmt.Sprintf("channel:%s", ch.Name),
|
||||
Priority: priority,
|
||||
Metadata: string(channelMetaBytes),
|
||||
ChannelID: &channelID,
|
||||
ReplyTo: replyTo,
|
||||
Attachments: attachments,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create channel message: %w", err)
|
||||
}
|
||||
|
||||
// 2. Deliver inbox notifications to other members.
|
||||
// 2. Deliver inbox notifications only to @mentioned members.
|
||||
recipientCount := 0
|
||||
for _, m := range members {
|
||||
if m.AgentName == fromAgent {
|
||||
if m.AgentName == fromAgent || !mentionedMembers[m.AgentName] {
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -537,12 +569,7 @@ func (s *Service) BroadcastMessage(ctx context.Context, channelID int64, fromAge
|
||||
"channel_id": channelID,
|
||||
"channel_name": ch.Name,
|
||||
"channel_message_id": channelMsg.ID,
|
||||
}
|
||||
if len(mentionedAgentsList) > 0 {
|
||||
inboxMetaObj["mentioned_agents"] = mentionedAgentsList
|
||||
}
|
||||
if mentionedMembers[m.AgentName] {
|
||||
inboxMetaObj["mention"] = true
|
||||
"mention": true,
|
||||
}
|
||||
inboxMetaBytes, _ := json.Marshal(inboxMetaObj)
|
||||
|
||||
@@ -552,7 +579,7 @@ func (s *Service) BroadcastMessage(ctx context.Context, channelID int64, fromAge
|
||||
Metadata: string(inboxMetaBytes),
|
||||
})
|
||||
if err != nil {
|
||||
s.logger.Error("failed to send channel notification",
|
||||
s.logger.Error("failed to send mention notification",
|
||||
"channel_id", channelID,
|
||||
"from", fromAgent,
|
||||
"to", m.AgentName,
|
||||
|
||||
@@ -533,7 +533,7 @@ func TestService_BroadcastMessage(t *testing.T) {
|
||||
ch, _ := svc.CreateChannel(ctx, CreateChannelRequest{Name: "alerts", Type: TypeStandard, CreatedBy: "agent-a"})
|
||||
|
||||
t.Run("broadcast creates channel message", func(t *testing.T) {
|
||||
msgs, err := svc.BroadcastMessage(ctx, ch.ID, "agent-a", "hello", 5, "")
|
||||
msgs, err := svc.BroadcastMessage(ctx, ch.ID, "agent-a", "hello", 5, "", nil, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("BroadcastMessage: %v", err)
|
||||
}
|
||||
@@ -565,30 +565,25 @@ func TestService_BroadcastMessage(t *testing.T) {
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("broadcast delivers inbox notifications", func(t *testing.T) {
|
||||
t.Run("broadcast without mentions sends no DMs", func(t *testing.T) {
|
||||
svc.JoinChannel(ctx, ch.ID, "agent-b")
|
||||
svc.JoinChannel(ctx, ch.ID, "agent-c")
|
||||
_, err := svc.BroadcastMessage(ctx, ch.ID, "agent-a", "multi-member test", 5, "")
|
||||
_, err := svc.BroadcastMessage(ctx, ch.ID, "agent-a", "no-dm-test", 5, "", nil, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("BroadcastMessage: %v", err)
|
||||
}
|
||||
|
||||
// agent-b should have an inbox notification
|
||||
// agent-b should NOT have an inbox notification (no @mention)
|
||||
inboxResult, _ := svc.msgService.ReadInbox(ctx, "agent-b", messaging.ReadOptions{IncludeRead: true})
|
||||
found := false
|
||||
for _, m := range inboxResult.Messages {
|
||||
if m.Body == "multi-member test" {
|
||||
found = true
|
||||
break
|
||||
if m.Body == "no-dm-test" {
|
||||
t.Error("agent-b should not receive inbox DM for non-mention broadcast")
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Error("agent-b did not receive inbox notification")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("sender does not receive own message", func(t *testing.T) {
|
||||
svc.BroadcastMessage(ctx, ch.ID, "agent-a", "no self-message", 5, "")
|
||||
svc.BroadcastMessage(ctx, ch.ID, "agent-a", "no self-message", 5, "", nil, nil)
|
||||
inboxResult, _ := svc.msgService.ReadInbox(ctx, "agent-a", messaging.ReadOptions{IncludeRead: true})
|
||||
for _, m := range inboxResult.Messages {
|
||||
if m.Body == "no self-message" {
|
||||
@@ -597,12 +592,42 @@ func TestService_BroadcastMessage(t *testing.T) {
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("non-member cannot broadcast", func(t *testing.T) {
|
||||
// agent-c is a member but let's test someone who isn't
|
||||
t.Run("non-member auto-joins public channel on broadcast", func(t *testing.T) {
|
||||
seedAgent(t, svc.store.(*SQLiteChannelStore).db, "outsider")
|
||||
_, err := svc.BroadcastMessage(ctx, ch.ID, "outsider", "unauthorized", 5, "")
|
||||
_, err := svc.BroadcastMessage(ctx, ch.ID, "outsider", "auto-joined", 5, "", nil, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("expected auto-join for public channel, got %v", err)
|
||||
}
|
||||
isMember, _ := svc.IsMember(ctx, ch.ID, "outsider")
|
||||
if !isMember {
|
||||
t.Error("outsider should be a member after auto-join")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("broadcast with reply_to", func(t *testing.T) {
|
||||
msgs, _ := svc.BroadcastMessage(ctx, ch.ID, "agent-a", "original", 5, "", nil, nil)
|
||||
original := msgs[0]
|
||||
|
||||
replies, err := svc.BroadcastMessage(ctx, ch.ID, "agent-b", "reply to original", 5, "", &original.ID, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("BroadcastMessage with reply_to: %v", err)
|
||||
}
|
||||
if replies[0].ReplyTo == nil || *replies[0].ReplyTo != original.ID {
|
||||
t.Errorf("reply_to = %v, want %d", replies[0].ReplyTo, original.ID)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("non-member cannot broadcast to private channel", func(t *testing.T) {
|
||||
privCh, err := svc.CreateChannel(ctx, CreateChannelRequest{
|
||||
Name: "private-test", Type: TypeStandard, IsPrivate: true, CreatedBy: "agent-a",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("create private channel: %v", err)
|
||||
}
|
||||
seedAgent(t, svc.store.(*SQLiteChannelStore).db, "outsider2")
|
||||
_, err = svc.BroadcastMessage(ctx, privCh.ID, "outsider2", "unauthorized", 5, "", nil, nil)
|
||||
if !errors.Is(err, ErrNotChannelMember) {
|
||||
t.Errorf("expected ErrNotChannelMember, got %v", err)
|
||||
t.Errorf("expected ErrNotChannelMember for private channel, got %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
@@ -619,7 +644,7 @@ func TestService_BroadcastMessage_Mentions(t *testing.T) {
|
||||
svc.JoinChannel(ctx, ch.ID, "agent-c")
|
||||
|
||||
t.Run("mentioned member gets mention flag in inbox", func(t *testing.T) {
|
||||
_, err := svc.BroadcastMessage(ctx, ch.ID, "agent-a", "hey @agent-b check this", 5, "")
|
||||
_, err := svc.BroadcastMessage(ctx, ch.ID, "agent-a", "hey @agent-b check this", 5, "", nil, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("BroadcastMessage: %v", err)
|
||||
}
|
||||
@@ -642,26 +667,17 @@ func TestService_BroadcastMessage_Mentions(t *testing.T) {
|
||||
t.Error("agent-b did not receive inbox notification")
|
||||
}
|
||||
|
||||
// agent-c was NOT mentioned — inbox notification should NOT have mention:true
|
||||
// agent-c was NOT mentioned — should NOT receive an inbox DM at all
|
||||
inboxResult, _ = svc.msgService.ReadInbox(ctx, "agent-c", messaging.ReadOptions{IncludeRead: true})
|
||||
for _, m := range inboxResult.Messages {
|
||||
if m.Body == "hey @agent-b check this" {
|
||||
var meta map[string]any
|
||||
json.Unmarshal(m.Metadata, &meta)
|
||||
if meta["mention"] == true {
|
||||
t.Error("agent-c should NOT have mention flag")
|
||||
}
|
||||
// But should still have mentioned_agents list
|
||||
if _, ok := meta["mentioned_agents"]; !ok {
|
||||
t.Error("agent-c metadata should have mentioned_agents list")
|
||||
}
|
||||
break
|
||||
t.Error("agent-c should not receive inbox DM when not @mentioned")
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("channel message metadata includes mentioned_agents", func(t *testing.T) {
|
||||
_, err := svc.BroadcastMessage(ctx, ch.ID, "agent-a", "cc @agent-b and @agent-c", 5, "")
|
||||
_, err := svc.BroadcastMessage(ctx, ch.ID, "agent-a", "cc @agent-b and @agent-c", 5, "", nil, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("BroadcastMessage: %v", err)
|
||||
}
|
||||
@@ -691,7 +707,7 @@ func TestService_BroadcastMessage_Mentions(t *testing.T) {
|
||||
})
|
||||
|
||||
t.Run("self-mention is excluded", func(t *testing.T) {
|
||||
_, err := svc.BroadcastMessage(ctx, ch.ID, "agent-a", "I am @agent-a and cc @agent-b", 5, "")
|
||||
_, err := svc.BroadcastMessage(ctx, ch.ID, "agent-a", "I am @agent-a and cc @agent-b", 5, "", nil, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("BroadcastMessage: %v", err)
|
||||
}
|
||||
@@ -716,7 +732,7 @@ func TestService_BroadcastMessage_Mentions(t *testing.T) {
|
||||
})
|
||||
|
||||
t.Run("no mentions produces no mention metadata", func(t *testing.T) {
|
||||
_, err := svc.BroadcastMessage(ctx, ch.ID, "agent-a", "just a normal message", 5, "")
|
||||
_, err := svc.BroadcastMessage(ctx, ch.ID, "agent-a", "just a normal message", 5, "", nil, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("BroadcastMessage: %v", err)
|
||||
}
|
||||
@@ -736,7 +752,7 @@ func TestService_BroadcastMessage_Mentions(t *testing.T) {
|
||||
|
||||
t.Run("non-member mention is ignored", func(t *testing.T) {
|
||||
seedAgent(t, svc.store.(*SQLiteChannelStore).db, "outsider")
|
||||
_, err := svc.BroadcastMessage(ctx, ch.ID, "agent-a", "hey @outsider and @agent-b", 5, "")
|
||||
_, err := svc.BroadcastMessage(ctx, ch.ID, "agent-a", "hey @outsider and @agent-b", 5, "", nil, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("BroadcastMessage: %v", err)
|
||||
}
|
||||
|
||||
@@ -83,9 +83,9 @@ func (s *SQLiteChannelStore) GetChannel(ctx context.Context, id int64) (*Channel
|
||||
var ch Channel
|
||||
var isPrivate, isSystem int
|
||||
err := s.db.QueryRowContext(ctx,
|
||||
`SELECT id, name, description, topic, type, is_private, is_system, created_by, created_at, updated_at
|
||||
`SELECT id, name, description, topic, type, is_private, is_system, created_by, workflow_enabled, auto_approve, stalemate_remind_after, stalemate_escalate_after, created_at, updated_at
|
||||
FROM channels WHERE id = ?`, id,
|
||||
).Scan(&ch.ID, &ch.Name, &ch.Description, &ch.Topic, &ch.Type, &isPrivate, &isSystem, &ch.CreatedBy, &ch.CreatedAt, &ch.UpdatedAt)
|
||||
).Scan(&ch.ID, &ch.Name, &ch.Description, &ch.Topic, &ch.Type, &isPrivate, &isSystem, &ch.CreatedBy, &ch.WorkflowEnabled, &ch.AutoApprove, &ch.StalemateRemindAfter, &ch.StalemateEscalateAfter, &ch.CreatedAt, &ch.UpdatedAt)
|
||||
if err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, ErrChannelNotFound
|
||||
@@ -102,9 +102,9 @@ func (s *SQLiteChannelStore) GetChannelByName(ctx context.Context, name string)
|
||||
var ch Channel
|
||||
var isPrivate, isSystem int
|
||||
err := s.db.QueryRowContext(ctx,
|
||||
`SELECT id, name, description, topic, type, is_private, is_system, created_by, created_at, updated_at
|
||||
`SELECT id, name, description, topic, type, is_private, is_system, created_by, workflow_enabled, auto_approve, stalemate_remind_after, stalemate_escalate_after, created_at, updated_at
|
||||
FROM channels WHERE LOWER(name) = LOWER(?)`, name,
|
||||
).Scan(&ch.ID, &ch.Name, &ch.Description, &ch.Topic, &ch.Type, &isPrivate, &isSystem, &ch.CreatedBy, &ch.CreatedAt, &ch.UpdatedAt)
|
||||
).Scan(&ch.ID, &ch.Name, &ch.Description, &ch.Topic, &ch.Type, &isPrivate, &isSystem, &ch.CreatedBy, &ch.WorkflowEnabled, &ch.AutoApprove, &ch.StalemateRemindAfter, &ch.StalemateEscalateAfter, &ch.CreatedAt, &ch.UpdatedAt)
|
||||
if err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, ErrChannelNotFound
|
||||
@@ -120,7 +120,7 @@ func (s *SQLiteChannelStore) GetChannelByName(ctx context.Context, name string)
|
||||
// is a member or has a pending invite.
|
||||
func (s *SQLiteChannelStore) ListChannels(ctx context.Context, agentName string) ([]*Channel, error) {
|
||||
rows, err := s.db.QueryContext(ctx,
|
||||
`SELECT DISTINCT c.id, c.name, c.description, c.topic, c.type, c.is_private, c.is_system, c.created_by, c.created_at, c.updated_at
|
||||
`SELECT DISTINCT c.id, c.name, c.description, c.topic, c.type, c.is_private, c.is_system, c.created_by, c.workflow_enabled, c.auto_approve, c.stalemate_remind_after, c.stalemate_escalate_after, c.created_at, c.updated_at
|
||||
FROM channels c
|
||||
WHERE c.is_private = 0
|
||||
OR EXISTS (SELECT 1 FROM channel_members cm WHERE cm.channel_id = c.id AND cm.agent_name = ?)
|
||||
@@ -137,7 +137,7 @@ func (s *SQLiteChannelStore) ListChannels(ctx context.Context, agentName string)
|
||||
for rows.Next() {
|
||||
var ch Channel
|
||||
var isPrivate, isSystem int
|
||||
if err := rows.Scan(&ch.ID, &ch.Name, &ch.Description, &ch.Topic, &ch.Type, &isPrivate, &isSystem, &ch.CreatedBy, &ch.CreatedAt, &ch.UpdatedAt); err != nil {
|
||||
if err := rows.Scan(&ch.ID, &ch.Name, &ch.Description, &ch.Topic, &ch.Type, &isPrivate, &isSystem, &ch.CreatedBy, &ch.WorkflowEnabled, &ch.AutoApprove, &ch.StalemateRemindAfter, &ch.StalemateEscalateAfter, &ch.CreatedAt, &ch.UpdatedAt); err != nil {
|
||||
return nil, fmt.Errorf("scan channel: %w", err)
|
||||
}
|
||||
ch.IsPrivate = isPrivate != 0
|
||||
@@ -415,6 +415,23 @@ func (s *SQLiteChannelStore) GetChannelSummaries(ctx context.Context, agentName
|
||||
return summaries, rows.Err()
|
||||
}
|
||||
|
||||
// UpdateChannelSettings updates the workflow-related settings for a channel.
|
||||
func (s *SQLiteChannelStore) UpdateChannelSettings(ctx context.Context, id int64, settings ChannelSettings) error {
|
||||
result, err := s.db.ExecContext(ctx,
|
||||
`UPDATE channels SET workflow_enabled = ?, auto_approve = ?, stalemate_remind_after = ?, stalemate_escalate_after = ?, updated_at = CURRENT_TIMESTAMP WHERE id = ?`,
|
||||
settings.WorkflowEnabled, settings.AutoApprove, settings.StalemateRemindAfter, settings.StalemateEscalateAfter, id,
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("update channel settings: %w", err)
|
||||
}
|
||||
rowsAffected, _ := result.RowsAffected()
|
||||
if rowsAffected == 0 {
|
||||
return ErrChannelNotFound
|
||||
}
|
||||
s.logger.Info("channel settings updated", "id", id)
|
||||
return nil
|
||||
}
|
||||
|
||||
// isUniqueConstraintError checks if an error is a SQLite unique constraint violation.
|
||||
func isUniqueConstraintError(err error) bool {
|
||||
return strings.Contains(err.Error(), "UNIQUE constraint failed")
|
||||
|
||||
+22
-10
@@ -25,16 +25,20 @@ const (
|
||||
|
||||
// Channel represents a named group communication space.
|
||||
type Channel struct {
|
||||
ID int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Description string `json:"description"`
|
||||
Topic string `json:"topic"`
|
||||
Type string `json:"type"`
|
||||
IsPrivate bool `json:"is_private"`
|
||||
IsSystem bool `json:"is_system"`
|
||||
CreatedBy string `json:"created_by"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
ID int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Description string `json:"description"`
|
||||
Topic string `json:"topic"`
|
||||
Type string `json:"type"`
|
||||
IsPrivate bool `json:"is_private"`
|
||||
IsSystem bool `json:"is_system"`
|
||||
CreatedBy string `json:"created_by"`
|
||||
WorkflowEnabled bool `json:"workflow_enabled"`
|
||||
AutoApprove bool `json:"auto_approve"`
|
||||
StalemateRemindAfter string `json:"stalemate_remind_after"`
|
||||
StalemateEscalateAfter string `json:"stalemate_escalate_after"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
// ChannelWithCount embeds Channel and adds a member count.
|
||||
@@ -107,6 +111,14 @@ type JoinChannelRequest struct {
|
||||
AgentName string `json:"agent_name"`
|
||||
}
|
||||
|
||||
// ChannelSettings holds workflow-related settings for a channel.
|
||||
type ChannelSettings struct {
|
||||
WorkflowEnabled bool `json:"workflow_enabled"`
|
||||
AutoApprove bool `json:"auto_approve"`
|
||||
StalemateRemindAfter string `json:"stalemate_remind_after"`
|
||||
StalemateEscalateAfter string `json:"stalemate_escalate_after"`
|
||||
}
|
||||
|
||||
// InviteRequest is the input for inviting an agent to a channel.
|
||||
type InviteRequest struct {
|
||||
ChannelID int64 `json:"channel_id"`
|
||||
|
||||
+202
-6
@@ -7,12 +7,14 @@ import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/synapbus/synapbus/internal/agents"
|
||||
"github.com/synapbus/synapbus/internal/attachments"
|
||||
"github.com/synapbus/synapbus/internal/channels"
|
||||
"github.com/synapbus/synapbus/internal/messaging"
|
||||
"github.com/synapbus/synapbus/internal/reactions"
|
||||
"github.com/synapbus/synapbus/internal/search"
|
||||
)
|
||||
|
||||
@@ -25,6 +27,7 @@ type ServiceBridge struct {
|
||||
swarmService *channels.SwarmService
|
||||
attachmentService *attachments.Service
|
||||
searchService *search.Service
|
||||
reactionService *reactions.Service
|
||||
agentName string
|
||||
}
|
||||
|
||||
@@ -36,6 +39,7 @@ func NewServiceBridge(
|
||||
swarmService *channels.SwarmService,
|
||||
attachmentService *attachments.Service,
|
||||
searchService *search.Service,
|
||||
reactionService *reactions.Service,
|
||||
agentName string,
|
||||
) *ServiceBridge {
|
||||
return &ServiceBridge{
|
||||
@@ -45,6 +49,7 @@ func NewServiceBridge(
|
||||
swarmService: swarmService,
|
||||
attachmentService: attachmentService,
|
||||
searchService: searchService,
|
||||
reactionService: reactionService,
|
||||
agentName: agentName,
|
||||
}
|
||||
}
|
||||
@@ -102,6 +107,16 @@ func (b *ServiceBridge) Call(ctx context.Context, actionName string, args map[st
|
||||
case "download_attachment":
|
||||
return b.callDownloadAttachment(ctx, args)
|
||||
|
||||
// --- Reactions ---
|
||||
case "react":
|
||||
return b.callReact(ctx, args)
|
||||
case "unreact":
|
||||
return b.callUnreact(ctx, args)
|
||||
case "get_reactions":
|
||||
return b.callGetReactions(ctx, args)
|
||||
case "list_by_state":
|
||||
return b.callListByState(ctx, args)
|
||||
|
||||
// --- DM send (also accessible via bridge for execute tool) ---
|
||||
case "send_message":
|
||||
return b.callSendMessage(ctx, args)
|
||||
@@ -132,12 +147,34 @@ func (b *ServiceBridge) callSendMessage(ctx context.Context, args map[string]any
|
||||
replyTo = &v
|
||||
}
|
||||
|
||||
var attachmentHashes []string
|
||||
if attVal, ok := args["attachments"]; ok {
|
||||
switch v := attVal.(type) {
|
||||
case string:
|
||||
for _, h := range strings.Split(v, ",") {
|
||||
h = strings.TrimSpace(h)
|
||||
if h != "" {
|
||||
attachmentHashes = append(attachmentHashes, h)
|
||||
}
|
||||
}
|
||||
case []any:
|
||||
for _, item := range v {
|
||||
if s, ok := item.(string); ok && s != "" {
|
||||
attachmentHashes = append(attachmentHashes, s)
|
||||
}
|
||||
}
|
||||
case []string:
|
||||
attachmentHashes = v
|
||||
}
|
||||
}
|
||||
|
||||
opts := messaging.SendOptions{
|
||||
Subject: getString(args, "subject", ""),
|
||||
Priority: getInt(args, "priority", 5),
|
||||
Metadata: getString(args, "metadata", ""),
|
||||
ChannelID: channelID,
|
||||
ReplyTo: replyTo,
|
||||
Subject: getString(args, "subject", ""),
|
||||
Priority: getInt(args, "priority", 5),
|
||||
Metadata: getString(args, "metadata", ""),
|
||||
ChannelID: channelID,
|
||||
ReplyTo: replyTo,
|
||||
Attachments: attachmentHashes,
|
||||
}
|
||||
|
||||
msg, err := b.msgService.SendMessage(ctx, b.agentName, to, body, opts)
|
||||
@@ -546,7 +583,36 @@ func (b *ServiceBridge) callSendChannelMessage(ctx context.Context, args map[str
|
||||
priority := getInt(args, "priority", 5)
|
||||
metadata := getString(args, "metadata", "")
|
||||
|
||||
messages, err := b.channelService.BroadcastMessage(ctx, channelID, b.agentName, body, priority, metadata)
|
||||
var replyTo *int64
|
||||
if v, ok := args["reply_to"]; ok {
|
||||
if f, ok := v.(float64); ok {
|
||||
r := int64(f)
|
||||
replyTo = &r
|
||||
}
|
||||
}
|
||||
|
||||
var attachmentHashes []string
|
||||
if attVal, ok := args["attachments"]; ok {
|
||||
switch v := attVal.(type) {
|
||||
case string:
|
||||
for _, h := range strings.Split(v, ",") {
|
||||
h = strings.TrimSpace(h)
|
||||
if h != "" {
|
||||
attachmentHashes = append(attachmentHashes, h)
|
||||
}
|
||||
}
|
||||
case []any:
|
||||
for _, item := range v {
|
||||
if s, ok := item.(string); ok && s != "" {
|
||||
attachmentHashes = append(attachmentHashes, s)
|
||||
}
|
||||
}
|
||||
case []string:
|
||||
attachmentHashes = v
|
||||
}
|
||||
}
|
||||
|
||||
messages, err := b.channelService.BroadcastMessage(ctx, channelID, b.agentName, body, priority, metadata, replyTo, attachmentHashes)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -868,6 +934,136 @@ func (b *ServiceBridge) callDownloadAttachment(ctx context.Context, args map[str
|
||||
}, nil
|
||||
}
|
||||
|
||||
// --- Reaction implementations ---
|
||||
|
||||
func (b *ServiceBridge) callReact(ctx context.Context, args map[string]any) (any, error) {
|
||||
if b.reactionService == nil {
|
||||
return nil, fmt.Errorf("reaction service not available")
|
||||
}
|
||||
|
||||
messageID := getInt(args, "message_id", 0)
|
||||
if messageID == 0 {
|
||||
return nil, fmt.Errorf("'message_id' parameter is required")
|
||||
}
|
||||
|
||||
reaction := getString(args, "reaction", "")
|
||||
if reaction == "" {
|
||||
return nil, fmt.Errorf("'reaction' parameter is required")
|
||||
}
|
||||
|
||||
var metadata json.RawMessage
|
||||
if metaStr := getString(args, "metadata", ""); metaStr != "" {
|
||||
if !json.Valid([]byte(metaStr)) {
|
||||
return nil, fmt.Errorf("metadata must be valid JSON")
|
||||
}
|
||||
metadata = json.RawMessage(metaStr)
|
||||
}
|
||||
|
||||
result, err := b.reactionService.Toggle(ctx, int64(messageID), b.agentName, reaction, metadata)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
resp := map[string]any{
|
||||
"action": result.Action,
|
||||
"message_id": messageID,
|
||||
"reaction": reaction,
|
||||
}
|
||||
if result.Reaction != nil {
|
||||
resp["id"] = result.Reaction.ID
|
||||
resp["created_at"] = result.Reaction.CreatedAt
|
||||
}
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
func (b *ServiceBridge) callUnreact(ctx context.Context, args map[string]any) (any, error) {
|
||||
if b.reactionService == nil {
|
||||
return nil, fmt.Errorf("reaction service not available")
|
||||
}
|
||||
|
||||
messageID := getInt(args, "message_id", 0)
|
||||
if messageID == 0 {
|
||||
return nil, fmt.Errorf("'message_id' parameter is required")
|
||||
}
|
||||
|
||||
reaction := getString(args, "reaction", "")
|
||||
if reaction == "" {
|
||||
return nil, fmt.Errorf("'reaction' parameter is required")
|
||||
}
|
||||
|
||||
if err := b.reactionService.Remove(ctx, int64(messageID), b.agentName, reaction); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return map[string]any{
|
||||
"message_id": messageID,
|
||||
"reaction": reaction,
|
||||
"status": "removed",
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (b *ServiceBridge) callGetReactions(ctx context.Context, args map[string]any) (any, error) {
|
||||
if b.reactionService == nil {
|
||||
return nil, fmt.Errorf("reaction service not available")
|
||||
}
|
||||
|
||||
messageID := getInt(args, "message_id", 0)
|
||||
if messageID == 0 {
|
||||
return nil, fmt.Errorf("'message_id' parameter is required")
|
||||
}
|
||||
|
||||
rxns, state, err := b.reactionService.GetReactions(ctx, int64(messageID))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return map[string]any{
|
||||
"reactions": rxns,
|
||||
"workflow_state": state,
|
||||
"count": len(rxns),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (b *ServiceBridge) callListByState(ctx context.Context, args map[string]any) (any, error) {
|
||||
if b.reactionService == nil {
|
||||
return nil, fmt.Errorf("reaction service not available")
|
||||
}
|
||||
if b.channelService == nil {
|
||||
return nil, fmt.Errorf("channel service not available")
|
||||
}
|
||||
|
||||
channelName := getString(args, "channel", "")
|
||||
if channelName == "" {
|
||||
return nil, fmt.Errorf("'channel' parameter is required")
|
||||
}
|
||||
|
||||
state := getString(args, "state", "")
|
||||
if state == "" {
|
||||
return nil, fmt.Errorf("'state' parameter is required")
|
||||
}
|
||||
|
||||
ch, err := b.channelService.GetChannelByName(ctx, channelName)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
messageIDs, err := b.reactionService.ListByState(ctx, ch.ID, state)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if messageIDs == nil {
|
||||
messageIDs = []int64{}
|
||||
}
|
||||
|
||||
return map[string]any{
|
||||
"message_ids": messageIDs,
|
||||
"count": len(messageIDs),
|
||||
"channel": channelName,
|
||||
"state": state,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// --- Helpers ---
|
||||
|
||||
// resolveChannelID resolves a channel ID from either channel_id or channel_name in args.
|
||||
|
||||
@@ -43,6 +43,7 @@ func newTestBridge(t *testing.T) (*ServiceBridge, *messaging.MessagingService, *
|
||||
swarmService,
|
||||
nil, // attachmentService
|
||||
nil, // searchService
|
||||
nil, // reactionService
|
||||
"agent-a",
|
||||
)
|
||||
return bridge, msgService, agentService, channelService
|
||||
@@ -185,7 +186,7 @@ func TestBridge_JoinChannel(t *testing.T) {
|
||||
bridge.agentService,
|
||||
bridge.channelService,
|
||||
bridge.swarmService,
|
||||
nil, nil,
|
||||
nil, nil, nil,
|
||||
"agent-b",
|
||||
)
|
||||
|
||||
|
||||
@@ -50,6 +50,7 @@ func newTestHybridWithChannels(t *testing.T) (*HybridToolRegistrar, *channels.Se
|
||||
nil, // swarmService
|
||||
nil, // attachmentService
|
||||
nil, // searchService
|
||||
nil, // reactionService
|
||||
jsPool,
|
||||
actionRegistry,
|
||||
actionIndex,
|
||||
|
||||
@@ -0,0 +1,503 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
mcplib "github.com/mark3labs/mcp-go/mcp"
|
||||
"github.com/mark3labs/mcp-go/server"
|
||||
|
||||
"github.com/synapbus/synapbus/internal/agents"
|
||||
"github.com/synapbus/synapbus/internal/channels"
|
||||
"github.com/synapbus/synapbus/internal/trace"
|
||||
)
|
||||
|
||||
// PromptRegistrar registers MCP prompts on the server.
|
||||
type PromptRegistrar struct {
|
||||
db *sql.DB
|
||||
agentService *agents.AgentService
|
||||
channelService *channels.Service
|
||||
traceStore trace.TraceStore
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
// NewPromptRegistrar creates a new prompt registrar.
|
||||
func NewPromptRegistrar(
|
||||
db *sql.DB,
|
||||
agentService *agents.AgentService,
|
||||
channelService *channels.Service,
|
||||
traceStore trace.TraceStore,
|
||||
) *PromptRegistrar {
|
||||
return &PromptRegistrar{
|
||||
db: db,
|
||||
agentService: agentService,
|
||||
channelService: channelService,
|
||||
traceStore: traceStore,
|
||||
logger: slog.Default().With("component", "mcp-prompts"),
|
||||
}
|
||||
}
|
||||
|
||||
// RegisterAllOnServer registers all 4 MCP prompts on an mcp-go MCPServer.
|
||||
func (p *PromptRegistrar) RegisterAllOnServer(s *server.MCPServer) {
|
||||
s.AddPrompt(p.dailyDigestPrompt(), p.handleDailyDigest)
|
||||
s.AddPrompt(p.agentHealthCheckPrompt(), p.handleAgentHealthCheck)
|
||||
s.AddPrompt(p.channelOverviewPrompt(), p.handleChannelOverview)
|
||||
s.AddPrompt(p.debugAgentPrompt(), p.handleDebugAgent)
|
||||
|
||||
p.logger.Info("MCP prompts registered", "count", 4)
|
||||
}
|
||||
|
||||
// --- Prompt Definitions ---
|
||||
|
||||
func (p *PromptRegistrar) dailyDigestPrompt() mcplib.Prompt {
|
||||
return mcplib.NewPrompt("daily-digest",
|
||||
mcplib.WithPromptDescription("Summary of SynapBus activity in the last 24 hours: message counts, active agents, busiest channels, and high-priority messages."),
|
||||
)
|
||||
}
|
||||
|
||||
func (p *PromptRegistrar) agentHealthCheckPrompt() mcplib.Prompt {
|
||||
return mcplib.NewPrompt("agent-health-check",
|
||||
mcplib.WithPromptDescription("Status overview of all registered agents: name, type, status, last activity, and pending DM count. Flags agents needing attention."),
|
||||
)
|
||||
}
|
||||
|
||||
func (p *PromptRegistrar) channelOverviewPrompt() mcplib.Prompt {
|
||||
return mcplib.NewPrompt("channel-overview",
|
||||
mcplib.WithPromptDescription("Overview of all channels: name, type, member count, and recent message activity."),
|
||||
)
|
||||
}
|
||||
|
||||
func (p *PromptRegistrar) debugAgentPrompt() mcplib.Prompt {
|
||||
return mcplib.NewPrompt("debug-agent",
|
||||
mcplib.WithPromptDescription("Detailed debug information for a specific agent: identity, recent traces, pending messages, and errors."),
|
||||
mcplib.WithArgument("agent_name",
|
||||
mcplib.ArgumentDescription("Name of the agent to debug"),
|
||||
mcplib.RequiredArgument(),
|
||||
),
|
||||
)
|
||||
}
|
||||
|
||||
// --- Prompt Handlers ---
|
||||
|
||||
func (p *PromptRegistrar) handleDailyDigest(ctx context.Context, req mcplib.GetPromptRequest) (*mcplib.GetPromptResult, error) {
|
||||
since := time.Now().Add(-24 * time.Hour)
|
||||
|
||||
var md strings.Builder
|
||||
md.WriteString("# SynapBus Daily Digest\n\n")
|
||||
md.WriteString(fmt.Sprintf("*Period: %s to %s*\n\n", since.Format(time.RFC3339), time.Now().Format(time.RFC3339)))
|
||||
|
||||
// Total messages sent in last 24h
|
||||
var totalMessages int
|
||||
err := p.db.QueryRowContext(ctx,
|
||||
`SELECT COUNT(*) FROM messages WHERE created_at >= ?`, since,
|
||||
).Scan(&totalMessages)
|
||||
if err != nil {
|
||||
md.WriteString(fmt.Sprintf("Error querying messages: %s\n\n", err))
|
||||
} else {
|
||||
md.WriteString(fmt.Sprintf("## Messages\n\n**Total messages sent**: %d\n\n", totalMessages))
|
||||
}
|
||||
|
||||
// Active agents (those who sent messages in last 24h)
|
||||
rows, err := p.db.QueryContext(ctx,
|
||||
`SELECT DISTINCT from_agent FROM messages WHERE created_at >= ? AND from_agent != '' ORDER BY from_agent`,
|
||||
since,
|
||||
)
|
||||
if err != nil {
|
||||
md.WriteString(fmt.Sprintf("Error querying active agents: %s\n\n", err))
|
||||
} else {
|
||||
var activeAgents []string
|
||||
for rows.Next() {
|
||||
var name string
|
||||
if err := rows.Scan(&name); err == nil {
|
||||
activeAgents = append(activeAgents, name)
|
||||
}
|
||||
}
|
||||
rows.Close()
|
||||
|
||||
md.WriteString("## Active Agents\n\n")
|
||||
if len(activeAgents) == 0 {
|
||||
md.WriteString("No agents were active in the last 24 hours.\n\n")
|
||||
} else {
|
||||
md.WriteString(fmt.Sprintf("**%d agents** sent messages: %s\n\n", len(activeAgents), strings.Join(activeAgents, ", ")))
|
||||
}
|
||||
}
|
||||
|
||||
// Top 3 busiest channels
|
||||
chanRows, err := p.db.QueryContext(ctx,
|
||||
`SELECT c.name, COUNT(m.id) as msg_count
|
||||
FROM messages m
|
||||
JOIN channels c ON m.channel_id = c.id
|
||||
WHERE m.created_at >= ? AND m.channel_id IS NOT NULL
|
||||
GROUP BY c.name
|
||||
ORDER BY msg_count DESC
|
||||
LIMIT 3`,
|
||||
since,
|
||||
)
|
||||
if err != nil {
|
||||
md.WriteString(fmt.Sprintf("Error querying channels: %s\n\n", err))
|
||||
} else {
|
||||
md.WriteString("## Top Channels\n\n")
|
||||
md.WriteString("| Channel | Messages |\n")
|
||||
md.WriteString("|---------|----------|\n")
|
||||
found := false
|
||||
for chanRows.Next() {
|
||||
var chName string
|
||||
var msgCount int
|
||||
if err := chanRows.Scan(&chName, &msgCount); err == nil {
|
||||
md.WriteString(fmt.Sprintf("| #%s | %d |\n", chName, msgCount))
|
||||
found = true
|
||||
}
|
||||
}
|
||||
chanRows.Close()
|
||||
if !found {
|
||||
md.WriteString("| (no channel activity) | - |\n")
|
||||
}
|
||||
md.WriteString("\n")
|
||||
}
|
||||
|
||||
// High-priority messages (priority >= 7)
|
||||
highRows, err := p.db.QueryContext(ctx,
|
||||
`SELECT id, from_agent, COALESCE(to_agent, ''), body, priority, created_at
|
||||
FROM messages
|
||||
WHERE created_at >= ? AND priority >= 7
|
||||
ORDER BY priority DESC, created_at DESC
|
||||
LIMIT 20`,
|
||||
since,
|
||||
)
|
||||
if err != nil {
|
||||
md.WriteString(fmt.Sprintf("Error querying high-priority messages: %s\n\n", err))
|
||||
} else {
|
||||
md.WriteString("## High-Priority Messages (>= 7)\n\n")
|
||||
var highPriorityFound bool
|
||||
for highRows.Next() {
|
||||
var id int64
|
||||
var from, to, body string
|
||||
var priority int
|
||||
var createdAt time.Time
|
||||
if err := highRows.Scan(&id, &from, &to, &body, &priority, &createdAt); err == nil {
|
||||
if !highPriorityFound {
|
||||
md.WriteString("| ID | From | To | Priority | Time | Body (truncated) |\n")
|
||||
md.WriteString("|----|------|----|----------|------|------------------|\n")
|
||||
highPriorityFound = true
|
||||
}
|
||||
if len(body) > 80 {
|
||||
body = body[:80] + "..."
|
||||
}
|
||||
// Escape pipe characters in body for markdown table
|
||||
body = strings.ReplaceAll(body, "|", "\\|")
|
||||
body = strings.ReplaceAll(body, "\n", " ")
|
||||
target := to
|
||||
if target == "" {
|
||||
target = "(channel)"
|
||||
}
|
||||
md.WriteString(fmt.Sprintf("| %d | %s | %s | %d | %s | %s |\n",
|
||||
id, from, target, priority, createdAt.Format("15:04"), body))
|
||||
}
|
||||
}
|
||||
highRows.Close()
|
||||
if !highPriorityFound {
|
||||
md.WriteString("No high-priority messages in the last 24 hours.\n")
|
||||
}
|
||||
md.WriteString("\n")
|
||||
}
|
||||
|
||||
return &mcplib.GetPromptResult{
|
||||
Description: "SynapBus activity summary for the last 24 hours",
|
||||
Messages: []mcplib.PromptMessage{
|
||||
{
|
||||
Role: mcplib.RoleUser,
|
||||
Content: mcplib.NewTextContent(md.String()),
|
||||
},
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (p *PromptRegistrar) handleAgentHealthCheck(ctx context.Context, req mcplib.GetPromptRequest) (*mcplib.GetPromptResult, error) {
|
||||
var md strings.Builder
|
||||
md.WriteString("# Agent Health Check\n\n")
|
||||
|
||||
// Get all active agents
|
||||
agentList, err := p.agentService.DiscoverAgents(ctx, "")
|
||||
if err != nil {
|
||||
return &mcplib.GetPromptResult{
|
||||
Description: "Agent health check",
|
||||
Messages: []mcplib.PromptMessage{
|
||||
{
|
||||
Role: mcplib.RoleUser,
|
||||
Content: mcplib.NewTextContent(fmt.Sprintf("Error listing agents: %s", err)),
|
||||
},
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
md.WriteString("| Agent | Type | Status | Last Activity | Pending DMs | Notes |\n")
|
||||
md.WriteString("|-------|------|--------|---------------|-------------|-------|\n")
|
||||
|
||||
for _, agent := range agentList {
|
||||
if agent.Name == "system" {
|
||||
continue
|
||||
}
|
||||
|
||||
// Get last activity from traces
|
||||
lastActivity := "unknown"
|
||||
if p.traceStore != nil {
|
||||
traces, err := p.traceStore.GetTraces(ctx, agent.Name, 1)
|
||||
if err == nil && len(traces) > 0 {
|
||||
lastActivity = traces[0].CreatedAt.Format("2006-01-02 15:04")
|
||||
}
|
||||
}
|
||||
|
||||
// Get pending DM count
|
||||
var pendingDMs int
|
||||
err := p.db.QueryRowContext(ctx,
|
||||
`SELECT COUNT(*) FROM messages WHERE to_agent = ? AND status = 'pending'`,
|
||||
agent.Name,
|
||||
).Scan(&pendingDMs)
|
||||
if err != nil {
|
||||
pendingDMs = -1
|
||||
}
|
||||
|
||||
notes := ""
|
||||
if pendingDMs > 10 {
|
||||
notes = "**NEEDS ATTENTION**"
|
||||
}
|
||||
|
||||
pendingStr := fmt.Sprintf("%d", pendingDMs)
|
||||
if pendingDMs < 0 {
|
||||
pendingStr = "error"
|
||||
}
|
||||
|
||||
md.WriteString(fmt.Sprintf("| %s | %s | %s | %s | %s | %s |\n",
|
||||
agent.Name, agent.Type, agent.Status, lastActivity, pendingStr, notes))
|
||||
}
|
||||
|
||||
md.WriteString(fmt.Sprintf("\n*Total agents: %d*\n", len(agentList)-1)) // exclude system
|
||||
|
||||
return &mcplib.GetPromptResult{
|
||||
Description: "Status of all registered agents",
|
||||
Messages: []mcplib.PromptMessage{
|
||||
{
|
||||
Role: mcplib.RoleUser,
|
||||
Content: mcplib.NewTextContent(md.String()),
|
||||
},
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (p *PromptRegistrar) handleChannelOverview(ctx context.Context, req mcplib.GetPromptRequest) (*mcplib.GetPromptResult, error) {
|
||||
since := time.Now().Add(-24 * time.Hour)
|
||||
|
||||
var md strings.Builder
|
||||
md.WriteString("# Channel Overview\n\n")
|
||||
|
||||
// Query all channels with member counts and recent message counts
|
||||
rows, err := p.db.QueryContext(ctx,
|
||||
`SELECT c.id, c.name, c.type, c.is_private,
|
||||
(SELECT COUNT(*) FROM channel_members cm WHERE cm.channel_id = c.id) as member_count,
|
||||
(SELECT COUNT(*) FROM messages m WHERE m.channel_id = c.id AND m.created_at >= ?) as msg_count_24h
|
||||
FROM channels c
|
||||
ORDER BY c.name`,
|
||||
since,
|
||||
)
|
||||
if err != nil {
|
||||
return &mcplib.GetPromptResult{
|
||||
Description: "Channel overview",
|
||||
Messages: []mcplib.PromptMessage{
|
||||
{
|
||||
Role: mcplib.RoleUser,
|
||||
Content: mcplib.NewTextContent(fmt.Sprintf("Error querying channels: %s", err)),
|
||||
},
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
md.WriteString("| Channel | Type | Private | Members | Messages (24h) |\n")
|
||||
md.WriteString("|---------|------|---------|---------|----------------|\n")
|
||||
|
||||
count := 0
|
||||
for rows.Next() {
|
||||
var id int64
|
||||
var name, chType string
|
||||
var isPrivate int
|
||||
var memberCount, msgCount int
|
||||
|
||||
if err := rows.Scan(&id, &name, &chType, &isPrivate, &memberCount, &msgCount); err != nil {
|
||||
continue
|
||||
}
|
||||
|
||||
privateStr := "no"
|
||||
if isPrivate != 0 {
|
||||
privateStr = "yes"
|
||||
}
|
||||
|
||||
md.WriteString(fmt.Sprintf("| #%s | %s | %s | %d | %d |\n",
|
||||
name, chType, privateStr, memberCount, msgCount))
|
||||
count++
|
||||
}
|
||||
|
||||
if count == 0 {
|
||||
md.WriteString("| (no channels) | - | - | - | - |\n")
|
||||
}
|
||||
|
||||
md.WriteString(fmt.Sprintf("\n*Total channels: %d*\n", count))
|
||||
|
||||
return &mcplib.GetPromptResult{
|
||||
Description: "Overview of all channels",
|
||||
Messages: []mcplib.PromptMessage{
|
||||
{
|
||||
Role: mcplib.RoleUser,
|
||||
Content: mcplib.NewTextContent(md.String()),
|
||||
},
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (p *PromptRegistrar) handleDebugAgent(ctx context.Context, req mcplib.GetPromptRequest) (*mcplib.GetPromptResult, error) {
|
||||
agentName := req.Params.Arguments["agent_name"]
|
||||
if agentName == "" {
|
||||
return &mcplib.GetPromptResult{
|
||||
Description: "Debug agent",
|
||||
Messages: []mcplib.PromptMessage{
|
||||
{
|
||||
Role: mcplib.RoleUser,
|
||||
Content: mcplib.NewTextContent("Error: agent_name argument is required"),
|
||||
},
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
var md strings.Builder
|
||||
md.WriteString(fmt.Sprintf("# Debug: Agent `%s`\n\n", agentName))
|
||||
|
||||
// Agent details
|
||||
agent, err := p.agentService.GetAgent(ctx, agentName)
|
||||
if err != nil {
|
||||
return &mcplib.GetPromptResult{
|
||||
Description: fmt.Sprintf("Debug info for agent %s", agentName),
|
||||
Messages: []mcplib.PromptMessage{
|
||||
{
|
||||
Role: mcplib.RoleUser,
|
||||
Content: mcplib.NewTextContent(fmt.Sprintf("Error: agent not found: %s", err)),
|
||||
},
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Resolve owner name
|
||||
ownerName := "unknown"
|
||||
if p.db != nil {
|
||||
var username sql.NullString
|
||||
_ = p.db.QueryRowContext(ctx,
|
||||
`SELECT username FROM users WHERE id = ?`, agent.OwnerID,
|
||||
).Scan(&username)
|
||||
if username.Valid {
|
||||
ownerName = username.String
|
||||
}
|
||||
}
|
||||
|
||||
md.WriteString("## Identity\n\n")
|
||||
md.WriteString(fmt.Sprintf("| Field | Value |\n"))
|
||||
md.WriteString(fmt.Sprintf("|-------|-------|\n"))
|
||||
md.WriteString(fmt.Sprintf("| Name | %s |\n", agent.Name))
|
||||
md.WriteString(fmt.Sprintf("| Display Name | %s |\n", agent.DisplayName))
|
||||
md.WriteString(fmt.Sprintf("| Type | %s |\n", agent.Type))
|
||||
md.WriteString(fmt.Sprintf("| Status | %s |\n", agent.Status))
|
||||
md.WriteString(fmt.Sprintf("| Owner | %s (ID: %d) |\n", ownerName, agent.OwnerID))
|
||||
md.WriteString(fmt.Sprintf("| Capabilities | `%s` |\n", string(agent.Capabilities)))
|
||||
md.WriteString(fmt.Sprintf("| Created | %s |\n", agent.CreatedAt.Format(time.RFC3339)))
|
||||
md.WriteString("\n")
|
||||
|
||||
// Pending messages count
|
||||
var pendingCount int
|
||||
err = p.db.QueryRowContext(ctx,
|
||||
`SELECT COUNT(*) FROM messages WHERE to_agent = ? AND status = 'pending'`,
|
||||
agentName,
|
||||
).Scan(&pendingCount)
|
||||
if err != nil {
|
||||
md.WriteString(fmt.Sprintf("Error querying pending messages: %s\n\n", err))
|
||||
} else {
|
||||
md.WriteString(fmt.Sprintf("## Pending Messages\n\n**Count**: %d\n\n", pendingCount))
|
||||
}
|
||||
|
||||
// Recent traces (last 10)
|
||||
if p.traceStore != nil {
|
||||
traces, err := p.traceStore.GetTraces(ctx, agentName, 10)
|
||||
if err != nil {
|
||||
md.WriteString(fmt.Sprintf("Error querying traces: %s\n\n", err))
|
||||
} else {
|
||||
md.WriteString("## Recent Traces (last 10)\n\n")
|
||||
if len(traces) == 0 {
|
||||
md.WriteString("No traces found.\n\n")
|
||||
} else {
|
||||
md.WriteString("| Time | Action | Details | Error |\n")
|
||||
md.WriteString("|------|--------|---------|-------|\n")
|
||||
for _, t := range traces {
|
||||
details := t.Details
|
||||
if len(details) > 80 {
|
||||
details = details[:80] + "..."
|
||||
}
|
||||
details = strings.ReplaceAll(details, "|", "\\|")
|
||||
details = strings.ReplaceAll(details, "\n", " ")
|
||||
|
||||
errStr := ""
|
||||
if t.Error.Valid {
|
||||
errStr = t.Error.String
|
||||
if len(errStr) > 60 {
|
||||
errStr = errStr[:60] + "..."
|
||||
}
|
||||
errStr = strings.ReplaceAll(errStr, "|", "\\|")
|
||||
}
|
||||
|
||||
md.WriteString(fmt.Sprintf("| %s | %s | %s | %s |\n",
|
||||
t.CreatedAt.Format("15:04:05"), t.Action, details, errStr))
|
||||
}
|
||||
md.WriteString("\n")
|
||||
}
|
||||
}
|
||||
|
||||
// Recent errors from traces
|
||||
errorTraces, err := p.traceStore.GetTraces(ctx, agentName, 50)
|
||||
if err == nil {
|
||||
md.WriteString("## Recent Errors\n\n")
|
||||
errorFound := false
|
||||
for _, t := range errorTraces {
|
||||
if t.Error.Valid && t.Error.String != "" {
|
||||
if !errorFound {
|
||||
md.WriteString("| Time | Action | Error |\n")
|
||||
md.WriteString("|------|--------|-------|\n")
|
||||
errorFound = true
|
||||
}
|
||||
errStr := t.Error.String
|
||||
if len(errStr) > 100 {
|
||||
errStr = errStr[:100] + "..."
|
||||
}
|
||||
errStr = strings.ReplaceAll(errStr, "|", "\\|")
|
||||
errStr = strings.ReplaceAll(errStr, "\n", " ")
|
||||
md.WriteString(fmt.Sprintf("| %s | %s | %s |\n",
|
||||
t.CreatedAt.Format("2006-01-02 15:04:05"), t.Action, errStr))
|
||||
}
|
||||
}
|
||||
if !errorFound {
|
||||
md.WriteString("No recent errors found.\n")
|
||||
}
|
||||
md.WriteString("\n")
|
||||
}
|
||||
} else {
|
||||
md.WriteString("## Traces\n\n*Trace store not available*\n\n")
|
||||
}
|
||||
|
||||
return &mcplib.GetPromptResult{
|
||||
Description: fmt.Sprintf("Debug information for agent %s", agentName),
|
||||
Messages: []mcplib.PromptMessage{
|
||||
{
|
||||
Role: mcplib.RoleUser,
|
||||
Content: mcplib.NewTextContent(md.String()),
|
||||
},
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
@@ -0,0 +1,339 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
mcplib "github.com/mark3labs/mcp-go/mcp"
|
||||
_ "modernc.org/sqlite"
|
||||
|
||||
"github.com/synapbus/synapbus/internal/agents"
|
||||
"github.com/synapbus/synapbus/internal/channels"
|
||||
"github.com/synapbus/synapbus/internal/messaging"
|
||||
"github.com/synapbus/synapbus/internal/trace"
|
||||
)
|
||||
|
||||
func newTestPromptRegistrar(t *testing.T) (*PromptRegistrar, *messaging.MessagingService, *agents.AgentService, *channels.Service) {
|
||||
t.Helper()
|
||||
db := newTestDB(t)
|
||||
|
||||
tracer := trace.NewTracer(db)
|
||||
t.Cleanup(func() { tracer.Close() })
|
||||
|
||||
msgStore := messaging.NewSQLiteMessageStore(db)
|
||||
msgService := messaging.NewMessagingService(msgStore, tracer)
|
||||
|
||||
agentStore := agents.NewSQLiteAgentStore(db)
|
||||
agentService := agents.NewAgentService(agentStore, tracer)
|
||||
|
||||
channelStore := channels.NewSQLiteChannelStore(db)
|
||||
channelService := channels.NewService(channelStore, msgService, tracer)
|
||||
|
||||
traceStore := trace.NewSQLiteTraceStore(db)
|
||||
|
||||
registrar := NewPromptRegistrar(db, agentService, channelService, traceStore)
|
||||
return registrar, msgService, agentService, channelService
|
||||
}
|
||||
|
||||
// extractPromptText extracts the text content from a GetPromptResult.
|
||||
func extractPromptText(t *testing.T, result *mcplib.GetPromptResult) string {
|
||||
t.Helper()
|
||||
if len(result.Messages) == 0 {
|
||||
t.Fatal("expected at least one prompt message")
|
||||
}
|
||||
tc, ok := result.Messages[0].Content.(mcplib.TextContent)
|
||||
if !ok {
|
||||
t.Fatalf("expected TextContent, got %T", result.Messages[0].Content)
|
||||
}
|
||||
return tc.Text
|
||||
}
|
||||
|
||||
func TestPrompt_DailyDigest(t *testing.T) {
|
||||
p, msgService, agentService, _ := newTestPromptRegistrar(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// Register agents
|
||||
agentService.Register(ctx, "agent-a", "Agent A", "ai", nil, 1)
|
||||
agentService.Register(ctx, "agent-b", "Agent B", "ai", nil, 1)
|
||||
|
||||
t.Run("empty system", func(t *testing.T) {
|
||||
req := mcplib.GetPromptRequest{}
|
||||
result, err := p.handleDailyDigest(ctx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleDailyDigest: %v", err)
|
||||
}
|
||||
|
||||
text := extractPromptText(t, result)
|
||||
if !strings.Contains(text, "Daily Digest") {
|
||||
t.Error("expected 'Daily Digest' in output")
|
||||
}
|
||||
if !strings.Contains(text, "Total messages sent") {
|
||||
t.Error("expected 'Total messages sent' in output")
|
||||
}
|
||||
if result.Messages[0].Role != mcplib.RoleUser {
|
||||
t.Errorf("expected role=user, got %s", result.Messages[0].Role)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("with messages", func(t *testing.T) {
|
||||
// Send some messages
|
||||
msgService.SendMessage(ctx, "agent-a", "agent-b", "hello", messaging.SendOptions{})
|
||||
msgService.SendMessage(ctx, "agent-b", "agent-a", "hi back", messaging.SendOptions{})
|
||||
msgService.SendMessage(ctx, "agent-a", "agent-b", "urgent!", messaging.SendOptions{Priority: 8})
|
||||
|
||||
req := mcplib.GetPromptRequest{}
|
||||
result, err := p.handleDailyDigest(ctx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleDailyDigest: %v", err)
|
||||
}
|
||||
|
||||
text := extractPromptText(t, result)
|
||||
if !strings.Contains(text, "agent-a") {
|
||||
t.Error("expected 'agent-a' in active agents")
|
||||
}
|
||||
if !strings.Contains(text, "agent-b") {
|
||||
t.Error("expected 'agent-b' in active agents")
|
||||
}
|
||||
// Should have the high priority message
|
||||
if !strings.Contains(text, "High-Priority") {
|
||||
t.Error("expected 'High-Priority' section in output")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestPrompt_AgentHealthCheck(t *testing.T) {
|
||||
p, msgService, agentService, _ := newTestPromptRegistrar(t)
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("no agents", func(t *testing.T) {
|
||||
req := mcplib.GetPromptRequest{}
|
||||
result, err := p.handleAgentHealthCheck(ctx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleAgentHealthCheck: %v", err)
|
||||
}
|
||||
|
||||
text := extractPromptText(t, result)
|
||||
if !strings.Contains(text, "Agent Health Check") {
|
||||
t.Error("expected 'Agent Health Check' header")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("with agents and pending DMs", func(t *testing.T) {
|
||||
agentService.Register(ctx, "healthy-agent", "Healthy", "ai", nil, 1)
|
||||
agentService.Register(ctx, "overloaded-agent", "Overloaded", "ai", nil, 1)
|
||||
|
||||
// Send 12 messages to overloaded-agent to trigger "NEEDS ATTENTION"
|
||||
for i := 0; i < 12; i++ {
|
||||
msgService.SendMessage(ctx, "healthy-agent", "overloaded-agent", "msg", messaging.SendOptions{})
|
||||
}
|
||||
|
||||
req := mcplib.GetPromptRequest{}
|
||||
result, err := p.handleAgentHealthCheck(ctx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleAgentHealthCheck: %v", err)
|
||||
}
|
||||
|
||||
text := extractPromptText(t, result)
|
||||
if !strings.Contains(text, "healthy-agent") {
|
||||
t.Error("expected 'healthy-agent' in output")
|
||||
}
|
||||
if !strings.Contains(text, "overloaded-agent") {
|
||||
t.Error("expected 'overloaded-agent' in output")
|
||||
}
|
||||
if !strings.Contains(text, "NEEDS ATTENTION") {
|
||||
t.Error("expected 'NEEDS ATTENTION' flag for overloaded agent")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestPrompt_ChannelOverview(t *testing.T) {
|
||||
p, _, agentService, channelService := newTestPromptRegistrar(t)
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("no channels", func(t *testing.T) {
|
||||
req := mcplib.GetPromptRequest{}
|
||||
result, err := p.handleChannelOverview(ctx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleChannelOverview: %v", err)
|
||||
}
|
||||
|
||||
text := extractPromptText(t, result)
|
||||
if !strings.Contains(text, "Channel Overview") {
|
||||
t.Error("expected 'Channel Overview' header")
|
||||
}
|
||||
if !strings.Contains(text, "Total channels: 0") {
|
||||
t.Error("expected 'Total channels: 0' with no channels")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("with channels", func(t *testing.T) {
|
||||
agentService.Register(ctx, "chan-creator", "Chan Creator", "ai", nil, 1)
|
||||
|
||||
channelService.CreateChannel(ctx, channels.CreateChannelRequest{
|
||||
Name: "general",
|
||||
Type: "standard",
|
||||
CreatedBy: "chan-creator",
|
||||
})
|
||||
channelService.CreateChannel(ctx, channels.CreateChannelRequest{
|
||||
Name: "private-ops",
|
||||
Type: "standard",
|
||||
IsPrivate: true,
|
||||
CreatedBy: "chan-creator",
|
||||
})
|
||||
|
||||
req := mcplib.GetPromptRequest{}
|
||||
result, err := p.handleChannelOverview(ctx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleChannelOverview: %v", err)
|
||||
}
|
||||
|
||||
text := extractPromptText(t, result)
|
||||
if !strings.Contains(text, "general") {
|
||||
t.Error("expected 'general' channel in output")
|
||||
}
|
||||
if !strings.Contains(text, "private-ops") {
|
||||
t.Error("expected 'private-ops' channel in output")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestPrompt_DebugAgent(t *testing.T) {
|
||||
p, msgService, agentService, _ := newTestPromptRegistrar(t)
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("missing agent_name", func(t *testing.T) {
|
||||
req := mcplib.GetPromptRequest{
|
||||
Params: mcplib.GetPromptParams{
|
||||
Name: "debug-agent",
|
||||
Arguments: map[string]string{},
|
||||
},
|
||||
}
|
||||
result, err := p.handleDebugAgent(ctx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleDebugAgent: %v", err)
|
||||
}
|
||||
|
||||
text := extractPromptText(t, result)
|
||||
if !strings.Contains(text, "agent_name argument is required") {
|
||||
t.Error("expected error message about missing agent_name")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("agent not found", func(t *testing.T) {
|
||||
req := mcplib.GetPromptRequest{
|
||||
Params: mcplib.GetPromptParams{
|
||||
Name: "debug-agent",
|
||||
Arguments: map[string]string{
|
||||
"agent_name": "nonexistent",
|
||||
},
|
||||
},
|
||||
}
|
||||
result, err := p.handleDebugAgent(ctx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleDebugAgent: %v", err)
|
||||
}
|
||||
|
||||
text := extractPromptText(t, result)
|
||||
if !strings.Contains(text, "agent not found") {
|
||||
t.Error("expected 'agent not found' error message")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("agent exists", func(t *testing.T) {
|
||||
agentService.Register(ctx, "debug-target", "Debug Target", "ai", nil, 1)
|
||||
|
||||
// Send some messages to create pending DMs
|
||||
agentService.Register(ctx, "sender-for-debug", "Sender", "ai", nil, 1)
|
||||
msgService.SendMessage(ctx, "sender-for-debug", "debug-target", "test message", messaging.SendOptions{})
|
||||
|
||||
// Wait briefly for traces to flush
|
||||
time.Sleep(200 * time.Millisecond)
|
||||
|
||||
req := mcplib.GetPromptRequest{
|
||||
Params: mcplib.GetPromptParams{
|
||||
Name: "debug-agent",
|
||||
Arguments: map[string]string{
|
||||
"agent_name": "debug-target",
|
||||
},
|
||||
},
|
||||
}
|
||||
result, err := p.handleDebugAgent(ctx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleDebugAgent: %v", err)
|
||||
}
|
||||
|
||||
text := extractPromptText(t, result)
|
||||
if !strings.Contains(text, "debug-target") {
|
||||
t.Error("expected 'debug-target' in output")
|
||||
}
|
||||
if !strings.Contains(text, "Identity") {
|
||||
t.Error("expected 'Identity' section in output")
|
||||
}
|
||||
if !strings.Contains(text, "Pending Messages") {
|
||||
t.Error("expected 'Pending Messages' section in output")
|
||||
}
|
||||
if !strings.Contains(text, "Recent Traces") {
|
||||
t.Error("expected 'Recent Traces' section in output")
|
||||
}
|
||||
if !strings.Contains(text, "Recent Errors") {
|
||||
t.Error("expected 'Recent Errors' section in output")
|
||||
}
|
||||
if !strings.Contains(text, "testowner") {
|
||||
t.Error("expected owner name 'testowner' in output")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestPrompt_Registration(t *testing.T) {
|
||||
p, _, _, _ := newTestPromptRegistrar(t)
|
||||
|
||||
// Verify prompt definitions
|
||||
t.Run("daily-digest prompt definition", func(t *testing.T) {
|
||||
prompt := p.dailyDigestPrompt()
|
||||
if prompt.Name != "daily-digest" {
|
||||
t.Errorf("name = %q, want daily-digest", prompt.Name)
|
||||
}
|
||||
if prompt.Description == "" {
|
||||
t.Error("expected non-empty description")
|
||||
}
|
||||
if len(prompt.Arguments) != 0 {
|
||||
t.Errorf("expected 0 arguments, got %d", len(prompt.Arguments))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("agent-health-check prompt definition", func(t *testing.T) {
|
||||
prompt := p.agentHealthCheckPrompt()
|
||||
if prompt.Name != "agent-health-check" {
|
||||
t.Errorf("name = %q, want agent-health-check", prompt.Name)
|
||||
}
|
||||
if len(prompt.Arguments) != 0 {
|
||||
t.Errorf("expected 0 arguments, got %d", len(prompt.Arguments))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("channel-overview prompt definition", func(t *testing.T) {
|
||||
prompt := p.channelOverviewPrompt()
|
||||
if prompt.Name != "channel-overview" {
|
||||
t.Errorf("name = %q, want channel-overview", prompt.Name)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("debug-agent prompt definition", func(t *testing.T) {
|
||||
prompt := p.debugAgentPrompt()
|
||||
if prompt.Name != "debug-agent" {
|
||||
t.Errorf("name = %q, want debug-agent", prompt.Name)
|
||||
}
|
||||
if len(prompt.Arguments) != 1 {
|
||||
t.Fatalf("expected 1 argument, got %d", len(prompt.Arguments))
|
||||
}
|
||||
arg := prompt.Arguments[0]
|
||||
if arg.Name != "agent_name" {
|
||||
t.Errorf("arg name = %q, want agent_name", arg.Name)
|
||||
}
|
||||
if !arg.Required {
|
||||
t.Error("expected agent_name to be required")
|
||||
}
|
||||
})
|
||||
}
|
||||
+10
-1
@@ -18,6 +18,7 @@ import (
|
||||
"github.com/synapbus/synapbus/internal/console"
|
||||
"github.com/synapbus/synapbus/internal/jsruntime"
|
||||
"github.com/synapbus/synapbus/internal/messaging"
|
||||
"github.com/synapbus/synapbus/internal/reactions"
|
||||
"github.com/synapbus/synapbus/internal/search"
|
||||
"github.com/synapbus/synapbus/internal/trace"
|
||||
)
|
||||
@@ -40,6 +41,7 @@ func NewMCPServer(
|
||||
swarmService *channels.SwarmService,
|
||||
attachmentService *attachments.Service,
|
||||
searchService *search.Service,
|
||||
reactionService *reactions.Service,
|
||||
consolePrinter *console.Printer,
|
||||
jsPool *jsruntime.Pool,
|
||||
actionRegistry *actions.Registry,
|
||||
@@ -141,6 +143,7 @@ func NewMCPServer(
|
||||
"SynapBus",
|
||||
"0.1.0",
|
||||
server.WithToolCapabilities(true),
|
||||
server.WithPromptCapabilities(true),
|
||||
server.WithHooks(hooks),
|
||||
)
|
||||
|
||||
@@ -152,6 +155,7 @@ func NewMCPServer(
|
||||
swarmService,
|
||||
attachmentService,
|
||||
searchService,
|
||||
reactionService,
|
||||
jsPool,
|
||||
actionRegistry,
|
||||
actionIndex,
|
||||
@@ -159,6 +163,11 @@ func NewMCPServer(
|
||||
)
|
||||
hybridRegistrar.RegisterAllOnServer(mcpSrv)
|
||||
|
||||
// Register the 4 MCP prompts
|
||||
traceStore := trace.NewSQLiteTraceStore(db)
|
||||
promptRegistrar := NewPromptRegistrar(db, agentService, channelService, traceStore)
|
||||
promptRegistrar.RegisterAllOnServer(mcpSrv)
|
||||
|
||||
// Create Streamable HTTP transport with context func for auth propagation
|
||||
httpServer := server.NewStreamableHTTPServer(mcpSrv,
|
||||
server.WithHTTPContextFunc(func(ctx context.Context, r *http.Request) context.Context {
|
||||
@@ -183,7 +192,7 @@ func NewMCPServer(
|
||||
console: consolePrinter,
|
||||
}
|
||||
|
||||
logger.Info("MCP server initialized (4 hybrid tools, streamable HTTP transport)")
|
||||
logger.Info("MCP server initialized (4 hybrid tools, 4 prompts, streamable HTTP transport)")
|
||||
return s
|
||||
}
|
||||
|
||||
|
||||
@@ -38,7 +38,7 @@ func newTestMCPServer(t *testing.T, con *console.Printer) (*MCPServer, *messagin
|
||||
actionRegistry := actions.NewRegistry()
|
||||
actionIndex := actions.NewIndex(actionRegistry.List())
|
||||
|
||||
srv := NewMCPServer(msgService, agentService, nil, nil, nil, nil, con, jsPool, actionRegistry, actionIndex, db)
|
||||
srv := NewMCPServer(msgService, agentService, nil, nil, nil, nil, nil, con, jsPool, actionRegistry, actionIndex, db)
|
||||
return srv, msgService, agentService
|
||||
}
|
||||
|
||||
@@ -133,7 +133,7 @@ func TestMCPToolCall_WithValidAPIKey(t *testing.T) {
|
||||
actionIndex := actions.NewIndex(actionRegistry.List())
|
||||
|
||||
// Create MCP server
|
||||
srv := NewMCPServer(msgService, agentService, nil, nil, nil, nil, nil, jsPool, actionRegistry, actionIndex, db)
|
||||
srv := NewMCPServer(msgService, agentService, nil, nil, nil, nil, nil, nil, jsPool, actionRegistry, actionIndex, db)
|
||||
|
||||
// Mount with auth middleware, just like main.go does
|
||||
mux := http.NewServeMux()
|
||||
@@ -188,7 +188,7 @@ func TestMCPToolCall_InvalidAPIKeyReturns401(t *testing.T) {
|
||||
actionRegistry := actions.NewRegistry()
|
||||
actionIndex := actions.NewIndex(actionRegistry.List())
|
||||
|
||||
srv := NewMCPServer(msgService, agentService, nil, nil, nil, nil, nil, jsPool, actionRegistry, actionIndex, db)
|
||||
srv := NewMCPServer(msgService, agentService, nil, nil, nil, nil, nil, nil, jsPool, actionRegistry, actionIndex, db)
|
||||
|
||||
mux := http.NewServeMux()
|
||||
handler := agents.OptionalAuthMiddlewareWithAPIKeys(agentService, apiKeyService)(srv.Handler())
|
||||
|
||||
@@ -18,6 +18,7 @@ import (
|
||||
"github.com/synapbus/synapbus/internal/channels"
|
||||
"github.com/synapbus/synapbus/internal/jsruntime"
|
||||
"github.com/synapbus/synapbus/internal/messaging"
|
||||
"github.com/synapbus/synapbus/internal/reactions"
|
||||
"github.com/synapbus/synapbus/internal/search"
|
||||
)
|
||||
|
||||
@@ -29,6 +30,7 @@ type HybridToolRegistrar struct {
|
||||
swarmService *channels.SwarmService
|
||||
attachmentService *attachments.Service
|
||||
searchService *search.Service
|
||||
reactionService *reactions.Service
|
||||
jsPool *jsruntime.Pool
|
||||
actionRegistry *actions.Registry
|
||||
actionIndex *actions.Index
|
||||
@@ -44,6 +46,7 @@ func NewHybridToolRegistrar(
|
||||
swarmService *channels.SwarmService,
|
||||
attachmentService *attachments.Service,
|
||||
searchService *search.Service,
|
||||
reactionService *reactions.Service,
|
||||
jsPool *jsruntime.Pool,
|
||||
actionRegistry *actions.Registry,
|
||||
actionIndex *actions.Index,
|
||||
@@ -56,6 +59,7 @@ func NewHybridToolRegistrar(
|
||||
swarmService: swarmService,
|
||||
attachmentService: attachmentService,
|
||||
searchService: searchService,
|
||||
reactionService: reactionService,
|
||||
jsPool: jsPool,
|
||||
actionRegistry: actionRegistry,
|
||||
actionIndex: actionIndex,
|
||||
@@ -84,14 +88,15 @@ func (h *HybridToolRegistrar) myStatusTool() mcplib.Tool {
|
||||
|
||||
func (h *HybridToolRegistrar) sendMessageTool() mcplib.Tool {
|
||||
return mcplib.NewTool("send_message",
|
||||
mcplib.WithDescription("Send a message to another agent (DM) or to a channel. Specify exactly one of 'to' (agent name for DM) or 'channel' (channel name or numeric ID)."),
|
||||
mcplib.WithDescription("Send a message to another agent (DM) or to a channel. Supports attachments — upload files first via the execute tool, then pass the returned hashes here. Specify exactly one of 'to' (agent name for DM) or 'channel' (channel name or numeric ID)."),
|
||||
mcplib.WithString("to", mcplib.Description("Recipient agent name for direct messages")),
|
||||
mcplib.WithString("channel", mcplib.Description("Channel name or numeric ID for channel messages")),
|
||||
mcplib.WithString("body", mcplib.Description("Message body text"), mcplib.Required()),
|
||||
mcplib.WithString("subject", mcplib.Description("Conversation subject (optional)")),
|
||||
mcplib.WithNumber("priority", mcplib.Description("Message priority (1-10, default 5)"), mcplib.Min(1), mcplib.Max(10)),
|
||||
mcplib.WithString("metadata", mcplib.Description("JSON metadata object (optional)")),
|
||||
mcplib.WithNumber("reply_to", mcplib.Description("ID of the message to reply to (optional, for threading)")),
|
||||
mcplib.WithNumber("reply_to", mcplib.Description("ID of the parent message to reply to. Creates a threaded reply. Always use reply_to when responding to a message that is itself a thread reply, to keep conversations organized.")),
|
||||
mcplib.WithString("attachments", mcplib.Description("Comma-separated list of attachment hashes to link to this message. Upload attachments first using the upload_attachment action via the execute tool.")),
|
||||
)
|
||||
}
|
||||
|
||||
@@ -323,6 +328,16 @@ func (h *HybridToolRegistrar) handleSendMessage(ctx context.Context, req mcplib.
|
||||
replyTo = &v
|
||||
}
|
||||
|
||||
var attachmentHashes []string
|
||||
if attStr := req.GetString("attachments", ""); attStr != "" {
|
||||
for _, h := range strings.Split(attStr, ",") {
|
||||
h = strings.TrimSpace(h)
|
||||
if h != "" {
|
||||
attachmentHashes = append(attachmentHashes, h)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Channel message path.
|
||||
if channel != "" {
|
||||
if h.channelService == nil {
|
||||
@@ -335,7 +350,7 @@ func (h *HybridToolRegistrar) handleSendMessage(ctx context.Context, req mcplib.
|
||||
return mcplib.NewToolResultError(fmt.Sprintf("send_message to channel failed: %s", err)), nil
|
||||
}
|
||||
|
||||
messages, err := h.channelService.BroadcastMessage(ctx, channelID, agentName, body, priority, metadataStr)
|
||||
messages, err := h.channelService.BroadcastMessage(ctx, channelID, agentName, body, priority, metadataStr, replyTo, attachmentHashes)
|
||||
if err != nil {
|
||||
return mcplib.NewToolResultError(fmt.Sprintf("send_message to channel failed: %s", err)), nil
|
||||
}
|
||||
@@ -345,19 +360,30 @@ func (h *HybridToolRegistrar) handleSendMessage(ctx context.Context, req mcplib.
|
||||
messageID = messages[0].ID
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
result := map[string]any{
|
||||
"channel_id": channelID,
|
||||
"message_id": messageID,
|
||||
"status": "sent",
|
||||
})
|
||||
}
|
||||
|
||||
// Enrich channel messages with attachment info.
|
||||
if len(messages) > 0 && len(attachmentHashes) > 0 {
|
||||
h.msgService.EnrichMessages(ctx, messages)
|
||||
if len(messages[0].Attachments) > 0 {
|
||||
result["attachments"] = messages[0].Attachments
|
||||
}
|
||||
}
|
||||
|
||||
return resultJSON(result)
|
||||
}
|
||||
|
||||
// DM path.
|
||||
opts := messaging.SendOptions{
|
||||
Subject: subject,
|
||||
Priority: priority,
|
||||
Metadata: metadataStr,
|
||||
ReplyTo: replyTo,
|
||||
Subject: subject,
|
||||
Priority: priority,
|
||||
Metadata: metadataStr,
|
||||
ReplyTo: replyTo,
|
||||
Attachments: attachmentHashes,
|
||||
}
|
||||
|
||||
msg, err := h.msgService.SendMessage(ctx, agentName, to, body, opts)
|
||||
@@ -365,11 +391,21 @@ func (h *HybridToolRegistrar) handleSendMessage(ctx context.Context, req mcplib.
|
||||
return mcplib.NewToolResultError(fmt.Sprintf("send_message failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
result := map[string]any{
|
||||
"message_id": msg.ID,
|
||||
"conversation_id": msg.ConversationID,
|
||||
"status": msg.Status,
|
||||
})
|
||||
}
|
||||
|
||||
// Enrich message with attachment info.
|
||||
if len(attachmentHashes) > 0 {
|
||||
h.msgService.EnrichMessages(ctx, []*messaging.Message{msg})
|
||||
if len(msg.Attachments) > 0 {
|
||||
result["attachments"] = msg.Attachments
|
||||
}
|
||||
}
|
||||
|
||||
return resultJSON(result)
|
||||
}
|
||||
|
||||
func (h *HybridToolRegistrar) handleSearch(ctx context.Context, req mcplib.CallToolRequest) (*mcplib.CallToolResult, error) {
|
||||
@@ -443,6 +479,7 @@ func (h *HybridToolRegistrar) handleExecute(ctx context.Context, req mcplib.Call
|
||||
h.swarmService,
|
||||
h.attachmentService,
|
||||
h.searchService,
|
||||
h.reactionService,
|
||||
agentName,
|
||||
)
|
||||
|
||||
|
||||
@@ -68,6 +68,7 @@ func newTestHybridRegistrar(t *testing.T) (*HybridToolRegistrar, *messaging.Mess
|
||||
nil, // swarmService
|
||||
nil, // attachmentService
|
||||
nil, // searchService
|
||||
nil, // reactionService
|
||||
jsPool,
|
||||
actionRegistry,
|
||||
actionIndex,
|
||||
|
||||
@@ -2,12 +2,13 @@ package messaging
|
||||
|
||||
// SendOptions configures message sending behavior.
|
||||
type SendOptions struct {
|
||||
Subject string `json:"subject,omitempty"`
|
||||
Priority int `json:"priority,omitempty"`
|
||||
Metadata string `json:"metadata,omitempty"`
|
||||
ChannelID *int64 `json:"channel_id,omitempty"`
|
||||
ConversationID *int64 `json:"conversation_id,omitempty"`
|
||||
ReplyTo *int64 `json:"reply_to,omitempty"`
|
||||
Subject string `json:"subject,omitempty"`
|
||||
Priority int `json:"priority,omitempty"`
|
||||
Metadata string `json:"metadata,omitempty"`
|
||||
ChannelID *int64 `json:"channel_id,omitempty"`
|
||||
ConversationID *int64 `json:"conversation_id,omitempty"`
|
||||
ReplyTo *int64 `json:"reply_to,omitempty"`
|
||||
Attachments []string `json:"attachments,omitempty"` // attachment hashes to link
|
||||
}
|
||||
|
||||
// ReadOptions configures inbox reading behavior.
|
||||
@@ -25,14 +26,18 @@ type ReadOptions struct {
|
||||
|
||||
// SearchOptions configures message search behavior.
|
||||
type SearchOptions struct {
|
||||
FromAgent string `json:"from_agent,omitempty"`
|
||||
ToAgent string `json:"to_agent,omitempty"`
|
||||
ChannelID *int64 `json:"channel_id,omitempty"`
|
||||
MinPriority int `json:"min_priority,omitempty"`
|
||||
Status string `json:"status,omitempty"`
|
||||
Limit int `json:"limit,omitempty"`
|
||||
Offset int `json:"offset,omitempty"`
|
||||
After string `json:"after,omitempty"`
|
||||
Before string `json:"before,omitempty"`
|
||||
Channel string `json:"channel,omitempty"`
|
||||
FromAgent string `json:"from_agent,omitempty"`
|
||||
ToAgent string `json:"to_agent,omitempty"`
|
||||
ChannelID *int64 `json:"channel_id,omitempty"`
|
||||
MinPriority int `json:"min_priority,omitempty"`
|
||||
Status string `json:"status,omitempty"`
|
||||
Limit int `json:"limit,omitempty"`
|
||||
Offset int `json:"offset,omitempty"`
|
||||
After string `json:"after,omitempty"`
|
||||
Before string `json:"before,omitempty"`
|
||||
Channel string `json:"channel,omitempty"`
|
||||
Channels []string `json:"channels,omitempty"` // include messages in these channels
|
||||
ExcludeChannels []string `json:"exclude_channels,omitempty"` // exclude messages in these channels
|
||||
Agents []string `json:"agents,omitempty"` // include messages from/to these agents
|
||||
ExcludeAgents []string `json:"exclude_agents,omitempty"` // exclude messages from/to these agents
|
||||
}
|
||||
|
||||
@@ -12,11 +12,40 @@ import (
|
||||
"github.com/synapbus/synapbus/internal/trace"
|
||||
)
|
||||
|
||||
// EmbeddingNotifier is called when messages are created or deleted so the
|
||||
// embedding pipeline can enqueue them without a direct import dependency.
|
||||
type EmbeddingNotifier interface {
|
||||
OnMessageCreated(ctx context.Context, messageID int64, body string)
|
||||
}
|
||||
|
||||
// MessageListener is notified after every message is persisted.
|
||||
// Implementations must not block — use goroutines for slow work.
|
||||
type MessageListener interface {
|
||||
OnMessageSent(ctx context.Context, msg *Message)
|
||||
}
|
||||
|
||||
// AttachmentLinker links attachment hashes to message IDs. This avoids
|
||||
// importing the attachments package directly. Set via SetAttachmentLinker.
|
||||
type AttachmentLinker interface {
|
||||
AttachToMessage(ctx context.Context, hash string, messageID int64) error
|
||||
GetByMessageID(ctx context.Context, messageID int64) ([]AttachmentInfo, error)
|
||||
}
|
||||
|
||||
// ReactionEnricher loads reactions for messages. This avoids importing the
|
||||
// reactions package directly. Set via SetReactionEnricher.
|
||||
type ReactionEnricher interface {
|
||||
GetByMessageIDs(ctx context.Context, messageIDs []int64) (map[int64][]ReactionInfo, error)
|
||||
}
|
||||
|
||||
// MessagingService provides business logic for messaging operations.
|
||||
type MessagingService struct {
|
||||
store MessageStore
|
||||
tracer *trace.Tracer
|
||||
dispatcher dispatcher.EventDispatcher
|
||||
embeddings EmbeddingNotifier
|
||||
attLinker AttachmentLinker
|
||||
rxEnricher ReactionEnricher
|
||||
listeners []MessageListener
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
@@ -34,6 +63,26 @@ func (s *MessagingService) SetDispatcher(d dispatcher.EventDispatcher) {
|
||||
s.dispatcher = d
|
||||
}
|
||||
|
||||
// SetEmbeddingNotifier sets the embedding pipeline callback for new messages.
|
||||
func (s *MessagingService) SetEmbeddingNotifier(n EmbeddingNotifier) {
|
||||
s.embeddings = n
|
||||
}
|
||||
|
||||
// SetAttachmentLinker sets the attachment linker for message-attachment binding.
|
||||
func (s *MessagingService) SetAttachmentLinker(l AttachmentLinker) {
|
||||
s.attLinker = l
|
||||
}
|
||||
|
||||
// SetReactionEnricher sets the reaction enricher for message reaction loading.
|
||||
func (s *MessagingService) SetReactionEnricher(e ReactionEnricher) {
|
||||
s.rxEnricher = e
|
||||
}
|
||||
|
||||
// AddMessageListener registers a listener that is notified after message creation.
|
||||
func (s *MessagingService) AddMessageListener(l MessageListener) {
|
||||
s.listeners = append(s.listeners, l)
|
||||
}
|
||||
|
||||
// SendMessage creates a message, auto-creating conversations as needed.
|
||||
func (s *MessagingService) SendMessage(ctx context.Context, from, to, body string, opts SendOptions) (*Message, error) {
|
||||
// Validate inputs
|
||||
@@ -121,6 +170,29 @@ func (s *MessagingService) SendMessage(ctx context.Context, from, to, body strin
|
||||
return nil, fmt.Errorf("insert message: %w", err)
|
||||
}
|
||||
|
||||
// Link attachments if provided.
|
||||
if s.attLinker != nil && len(opts.Attachments) > 0 {
|
||||
for _, hash := range opts.Attachments {
|
||||
if err := s.attLinker.AttachToMessage(ctx, hash, msg.ID); err != nil {
|
||||
s.logger.Error("failed to link attachment",
|
||||
"hash", hash,
|
||||
"message_id", msg.ID,
|
||||
"error", err,
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Enqueue for embedding (async, best-effort)
|
||||
if s.embeddings != nil {
|
||||
s.embeddings.OnMessageCreated(ctx, msg.ID, msg.Body)
|
||||
}
|
||||
|
||||
// Notify listeners (SSE, etc.)
|
||||
for _, l := range s.listeners {
|
||||
l.OnMessageSent(ctx, msg)
|
||||
}
|
||||
|
||||
s.logger.Info("message sent",
|
||||
"from", from,
|
||||
"to", to,
|
||||
@@ -489,6 +561,76 @@ func (s *MessagingService) GetConversationIDsForDM(ctx context.Context, agentNam
|
||||
return s.store.GetConversationIDsForDM(ctx, agentNames, peerAgent, lastMessageID)
|
||||
}
|
||||
|
||||
// EnrichMessages populates ReplyCount and Attachments on a slice of messages.
|
||||
func (s *MessagingService) EnrichMessages(ctx context.Context, msgs []*Message) {
|
||||
if len(msgs) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
ids := make([]int64, len(msgs))
|
||||
for i, m := range msgs {
|
||||
ids[i] = m.ID
|
||||
}
|
||||
|
||||
// Batch-load reply counts.
|
||||
counts, err := s.store.GetReplyCounts(ctx, ids)
|
||||
if err != nil {
|
||||
s.logger.Error("failed to load reply counts", "error", err)
|
||||
} else {
|
||||
for _, m := range msgs {
|
||||
if c, ok := counts[m.ID]; ok {
|
||||
m.ReplyCount = c
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Batch-load attachments.
|
||||
if s.attLinker != nil {
|
||||
for _, m := range msgs {
|
||||
atts, err := s.attLinker.GetByMessageID(ctx, m.ID)
|
||||
if err != nil {
|
||||
s.logger.Error("failed to load attachments", "message_id", m.ID, "error", err)
|
||||
continue
|
||||
}
|
||||
if len(atts) > 0 {
|
||||
m.Attachments = atts
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Batch-load reactions and derive workflow state.
|
||||
if s.rxEnricher != nil {
|
||||
rxMap, err := s.rxEnricher.GetByMessageIDs(ctx, ids)
|
||||
if err != nil {
|
||||
s.logger.Error("failed to load reactions", "error", err)
|
||||
} else {
|
||||
for _, m := range msgs {
|
||||
if rxs, ok := rxMap[m.ID]; ok && len(rxs) > 0 {
|
||||
m.Reactions = rxs
|
||||
// Derive workflow state from reactions
|
||||
highestPriority := 0
|
||||
priorities := map[string]int{
|
||||
"approve": 2, "in_progress": 3, "reject": 4, "done": 5, "published": 6,
|
||||
}
|
||||
states := map[string]string{
|
||||
"approve": "approved", "in_progress": "in_progress", "reject": "rejected", "done": "done", "published": "published",
|
||||
}
|
||||
for _, rx := range rxs {
|
||||
if p, ok := priorities[rx.Reaction]; ok && p > highestPriority {
|
||||
highestPriority = p
|
||||
m.WorkflowState = states[rx.Reaction]
|
||||
}
|
||||
}
|
||||
}
|
||||
// Default to "proposed" for channel messages with no reactions
|
||||
if m.WorkflowState == "" && m.ChannelID != nil {
|
||||
m.WorkflowState = "proposed"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// GetConversation returns a conversation and its messages.
|
||||
func (s *MessagingService) GetConversation(ctx context.Context, id int64) (*Conversation, []*Message, error) {
|
||||
conv, err := s.store.GetConversation(ctx, id)
|
||||
|
||||
@@ -728,5 +728,96 @@ func TestMessagingService_ReadInbox_DateFiltering(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
// mockAttachmentLinker is a test double for the AttachmentLinker interface.
|
||||
type mockAttachmentLinker struct {
|
||||
attachments map[int64][]AttachmentInfo
|
||||
}
|
||||
|
||||
func (m *mockAttachmentLinker) AttachToMessage(_ context.Context, _ string, _ int64) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *mockAttachmentLinker) GetByMessageID(_ context.Context, messageID int64) ([]AttachmentInfo, error) {
|
||||
return m.attachments[messageID], nil
|
||||
}
|
||||
|
||||
func TestMessagingService_EnrichMessages(t *testing.T) {
|
||||
svc, _ := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// Send a parent message and replies to it.
|
||||
parent, err := svc.SendMessage(ctx, "sender", "receiver", "parent message", SendOptions{Subject: "Enrich Test"})
|
||||
if err != nil {
|
||||
t.Fatalf("SendMessage (parent): %v", err)
|
||||
}
|
||||
|
||||
replyTo := parent.ID
|
||||
for i := 0; i < 3; i++ {
|
||||
_, err := svc.SendMessage(ctx, "receiver", "sender", "reply", SendOptions{
|
||||
Subject: "Enrich Test",
|
||||
ReplyTo: &replyTo,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("SendMessage (reply %d): %v", i, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Send a message with no replies.
|
||||
noReplies, err := svc.SendMessage(ctx, "sender", "receiver", "standalone", SendOptions{Subject: "Enrich Standalone"})
|
||||
if err != nil {
|
||||
t.Fatalf("SendMessage (standalone): %v", err)
|
||||
}
|
||||
|
||||
t.Run("reply counts populated", func(t *testing.T) {
|
||||
msgs := []*Message{parent, noReplies}
|
||||
svc.EnrichMessages(ctx, msgs)
|
||||
|
||||
if parent.ReplyCount != 3 {
|
||||
t.Errorf("parent ReplyCount = %d, want 3", parent.ReplyCount)
|
||||
}
|
||||
if noReplies.ReplyCount != 0 {
|
||||
t.Errorf("noReplies ReplyCount = %d, want 0", noReplies.ReplyCount)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("attachments populated when linker set", func(t *testing.T) {
|
||||
linker := &mockAttachmentLinker{
|
||||
attachments: map[int64][]AttachmentInfo{
|
||||
parent.ID: {
|
||||
{Hash: "abc123", OriginalFilename: "photo.png", Size: 1024, MIMEType: "image/png", IsImage: true},
|
||||
},
|
||||
},
|
||||
}
|
||||
svc.SetAttachmentLinker(linker)
|
||||
|
||||
// Reset enrichment state.
|
||||
parent.ReplyCount = 0
|
||||
parent.Attachments = nil
|
||||
noReplies.ReplyCount = 0
|
||||
noReplies.Attachments = nil
|
||||
|
||||
msgs := []*Message{parent, noReplies}
|
||||
svc.EnrichMessages(ctx, msgs)
|
||||
|
||||
if parent.ReplyCount != 3 {
|
||||
t.Errorf("parent ReplyCount = %d, want 3", parent.ReplyCount)
|
||||
}
|
||||
if len(parent.Attachments) != 1 {
|
||||
t.Fatalf("parent Attachments count = %d, want 1", len(parent.Attachments))
|
||||
}
|
||||
if parent.Attachments[0].Hash != "abc123" {
|
||||
t.Errorf("attachment hash = %s, want abc123", parent.Attachments[0].Hash)
|
||||
}
|
||||
if noReplies.Attachments != nil {
|
||||
t.Errorf("noReplies Attachments should be nil, got %v", noReplies.Attachments)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("empty slice is a no-op", func(t *testing.T) {
|
||||
svc.EnrichMessages(ctx, []*Message{})
|
||||
// Should not panic or error.
|
||||
})
|
||||
}
|
||||
|
||||
// suppress unused import warning for storage package
|
||||
var _ = storage.RunMigrations
|
||||
|
||||
@@ -0,0 +1,471 @@
|
||||
package messaging
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"os"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// StalemateConfig holds stalemate detection settings.
|
||||
type StalemateConfig struct {
|
||||
// ProcessingTimeout is how long a message can stay in "processing" before auto-fail (default 24h).
|
||||
ProcessingTimeout time.Duration
|
||||
// ReminderAfter is how long a pending DM waits before a system reminder is sent (default 4h).
|
||||
ReminderAfter time.Duration
|
||||
// EscalateAfter is how long a pending DM waits before escalation to #approvals (default 48h).
|
||||
EscalateAfter time.Duration
|
||||
// Interval is how often the worker checks for stale messages (default 15m).
|
||||
Interval time.Duration
|
||||
}
|
||||
|
||||
// DefaultStalemateConfig returns the default stalemate configuration.
|
||||
func DefaultStalemateConfig() StalemateConfig {
|
||||
return StalemateConfig{
|
||||
ProcessingTimeout: 24 * time.Hour,
|
||||
ReminderAfter: 4 * time.Hour,
|
||||
EscalateAfter: 48 * time.Hour,
|
||||
Interval: 15 * time.Minute,
|
||||
}
|
||||
}
|
||||
|
||||
// parseDurationWithDays parses a duration string supporting "Nd" format for days
|
||||
// in addition to standard Go duration formats.
|
||||
func parseDurationWithDays(s string) (time.Duration, error) {
|
||||
s = strings.TrimSpace(s)
|
||||
if s == "" {
|
||||
return 0, fmt.Errorf("empty duration string")
|
||||
}
|
||||
|
||||
// Try "Nd" format (days)
|
||||
if strings.HasSuffix(s, "d") {
|
||||
days, err := strconv.Atoi(strings.TrimSuffix(s, "d"))
|
||||
if err == nil && days > 0 {
|
||||
return time.Duration(days) * 24 * time.Hour, nil
|
||||
}
|
||||
}
|
||||
|
||||
// Try standard Go duration
|
||||
return time.ParseDuration(s)
|
||||
}
|
||||
|
||||
// ParseStalemateConfig reads stalemate configuration from environment variables.
|
||||
func ParseStalemateConfig() StalemateConfig {
|
||||
cfg := DefaultStalemateConfig()
|
||||
|
||||
if v := os.Getenv("SYNAPBUS_STALEMATE_PROCESSING_TIMEOUT"); v != "" {
|
||||
if d, err := parseDurationWithDays(v); err == nil && d > 0 {
|
||||
cfg.ProcessingTimeout = d
|
||||
}
|
||||
}
|
||||
if v := os.Getenv("SYNAPBUS_STALEMATE_REMINDER_AFTER"); v != "" {
|
||||
if d, err := parseDurationWithDays(v); err == nil && d > 0 {
|
||||
cfg.ReminderAfter = d
|
||||
}
|
||||
}
|
||||
if v := os.Getenv("SYNAPBUS_STALEMATE_ESCALATE_AFTER"); v != "" {
|
||||
if d, err := parseDurationWithDays(v); err == nil && d > 0 {
|
||||
cfg.EscalateAfter = d
|
||||
}
|
||||
}
|
||||
if v := os.Getenv("SYNAPBUS_STALEMATE_INTERVAL"); v != "" {
|
||||
if d, err := parseDurationWithDays(v); err == nil && d > 0 {
|
||||
cfg.Interval = d
|
||||
}
|
||||
}
|
||||
|
||||
return cfg
|
||||
}
|
||||
|
||||
// ChannelLookup provides channel lookup by name without importing the channels package.
|
||||
type ChannelLookup interface {
|
||||
// GetChannelIDByName returns a channel ID by name, or 0 if not found.
|
||||
GetChannelIDByName(ctx context.Context, name string) (int64, error)
|
||||
}
|
||||
|
||||
// StalemateWorker periodically checks for and handles stale messages.
|
||||
type StalemateWorker struct {
|
||||
db *sql.DB
|
||||
msgService *MessagingService
|
||||
channelLookup ChannelLookup
|
||||
config StalemateConfig
|
||||
logger *slog.Logger
|
||||
done chan struct{}
|
||||
wg sync.WaitGroup
|
||||
}
|
||||
|
||||
// NewStalemateWorker creates a new stalemate detection worker.
|
||||
func NewStalemateWorker(db *sql.DB, msgService *MessagingService, channelLookup ChannelLookup, config StalemateConfig) *StalemateWorker {
|
||||
return &StalemateWorker{
|
||||
db: db,
|
||||
msgService: msgService,
|
||||
channelLookup: channelLookup,
|
||||
config: config,
|
||||
logger: slog.Default().With("component", "stalemate-worker"),
|
||||
done: make(chan struct{}),
|
||||
}
|
||||
}
|
||||
|
||||
// Start begins the background stalemate check loop.
|
||||
func (w *StalemateWorker) Start() {
|
||||
w.wg.Add(1)
|
||||
go func() {
|
||||
defer w.wg.Done()
|
||||
w.logger.Info("stalemate worker started",
|
||||
"interval", w.config.Interval.String(),
|
||||
"processing_timeout", w.config.ProcessingTimeout.String(),
|
||||
"reminder_after", w.config.ReminderAfter.String(),
|
||||
"escalate_after", w.config.EscalateAfter.String(),
|
||||
)
|
||||
|
||||
ticker := time.NewTicker(w.config.Interval)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ticker.C:
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 60*time.Second)
|
||||
w.checkStaleMessages(ctx)
|
||||
cancel()
|
||||
case <-w.done:
|
||||
w.logger.Info("stalemate worker stopped")
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
// Stop stops the stalemate worker and waits for it to finish.
|
||||
func (w *StalemateWorker) Stop() {
|
||||
close(w.done)
|
||||
w.wg.Wait()
|
||||
}
|
||||
|
||||
// checkStaleMessages runs all stalemate checks.
|
||||
func (w *StalemateWorker) checkStaleMessages(ctx context.Context) {
|
||||
failed := w.failTimedOutProcessing(ctx)
|
||||
reminded := w.sendPendingReminders(ctx)
|
||||
escalated := w.escalatePendingMessages(ctx)
|
||||
|
||||
if failed > 0 || reminded > 0 || escalated > 0 {
|
||||
w.logger.Info("stalemate check complete",
|
||||
"auto_failed", failed,
|
||||
"reminders_sent", reminded,
|
||||
"escalations_sent", escalated,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
// staleDM represents a stale direct message found by the worker.
|
||||
type staleDM struct {
|
||||
ID int64
|
||||
FromAgent string
|
||||
ToAgent string
|
||||
Body string
|
||||
ClaimedAt *time.Time
|
||||
ClaimedBy string
|
||||
CreatedAt time.Time
|
||||
}
|
||||
|
||||
// failTimedOutProcessing auto-fails DMs in "processing" status that have exceeded the timeout.
|
||||
func (w *StalemateWorker) failTimedOutProcessing(ctx context.Context) int64 {
|
||||
cutoff := time.Now().Add(-w.config.ProcessingTimeout)
|
||||
|
||||
rows, err := w.db.QueryContext(ctx,
|
||||
`SELECT id, from_agent, to_agent, body, claimed_at, claimed_by
|
||||
FROM messages
|
||||
WHERE status = 'processing'
|
||||
AND to_agent IS NOT NULL
|
||||
AND to_agent != ''
|
||||
AND to_agent != 'system'
|
||||
AND claimed_at < ?`,
|
||||
cutoff,
|
||||
)
|
||||
if err != nil {
|
||||
w.logger.Error("query timed-out processing messages failed", "error", err)
|
||||
return 0
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var stale []staleDM
|
||||
for rows.Next() {
|
||||
var dm staleDM
|
||||
var claimedAt sql.NullTime
|
||||
var claimedBy sql.NullString
|
||||
if err := rows.Scan(&dm.ID, &dm.FromAgent, &dm.ToAgent, &dm.Body, &claimedAt, &claimedBy); err != nil {
|
||||
w.logger.Error("scan timed-out message failed", "error", err)
|
||||
continue
|
||||
}
|
||||
if claimedAt.Valid {
|
||||
dm.ClaimedAt = &claimedAt.Time
|
||||
}
|
||||
if claimedBy.Valid {
|
||||
dm.ClaimedBy = claimedBy.String
|
||||
}
|
||||
stale = append(stale, dm)
|
||||
}
|
||||
|
||||
count := int64(0)
|
||||
for _, dm := range stale {
|
||||
metadata := map[string]any{"error": "claim timeout exceeded"}
|
||||
metaBytes, _ := json.Marshal(metadata)
|
||||
|
||||
// Update directly via DB since the store's UpdateMessageStatus requires the claiming agent
|
||||
_, err := w.db.ExecContext(ctx,
|
||||
`UPDATE messages SET status = ?, metadata = ?, updated_at = CURRENT_TIMESTAMP
|
||||
WHERE id = ? AND status = 'processing'`,
|
||||
StatusFailed, string(metaBytes), dm.ID,
|
||||
)
|
||||
if err != nil {
|
||||
w.logger.Error("auto-fail message failed",
|
||||
"message_id", dm.ID,
|
||||
"error", err,
|
||||
)
|
||||
continue
|
||||
}
|
||||
w.logger.Info("auto-failed stale processing message",
|
||||
"message_id", dm.ID,
|
||||
"from_agent", dm.FromAgent,
|
||||
"to_agent", dm.ToAgent,
|
||||
"claimed_by", dm.ClaimedBy,
|
||||
)
|
||||
count++
|
||||
}
|
||||
return count
|
||||
}
|
||||
|
||||
// sendPendingReminders sends system DM reminders for pending messages older than ReminderAfter.
|
||||
func (w *StalemateWorker) sendPendingReminders(ctx context.Context) int64 {
|
||||
cutoff := time.Now().Add(-w.config.ReminderAfter)
|
||||
|
||||
rows, err := w.db.QueryContext(ctx,
|
||||
`SELECT id, from_agent, to_agent, body, created_at
|
||||
FROM messages
|
||||
WHERE status = 'pending'
|
||||
AND to_agent IS NOT NULL
|
||||
AND to_agent != ''
|
||||
AND from_agent != 'system'
|
||||
AND to_agent != 'system'
|
||||
AND created_at < ?`,
|
||||
cutoff,
|
||||
)
|
||||
if err != nil {
|
||||
w.logger.Error("query pending reminder candidates failed", "error", err)
|
||||
return 0
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
type pendingMsg struct {
|
||||
ID int64
|
||||
FromAgent string
|
||||
ToAgent string
|
||||
Body string
|
||||
CreatedAt time.Time
|
||||
}
|
||||
|
||||
var pending []pendingMsg
|
||||
for rows.Next() {
|
||||
var pm pendingMsg
|
||||
if err := rows.Scan(&pm.ID, &pm.FromAgent, &pm.ToAgent, &pm.Body, &pm.CreatedAt); err != nil {
|
||||
w.logger.Error("scan pending message failed", "error", err)
|
||||
continue
|
||||
}
|
||||
pending = append(pending, pm)
|
||||
}
|
||||
|
||||
count := int64(0)
|
||||
for _, pm := range pending {
|
||||
// Check if a reminder already exists for this message
|
||||
if w.reminderExists(ctx, pm.ID, pm.ToAgent) {
|
||||
continue
|
||||
}
|
||||
|
||||
age := formatAge(time.Since(pm.CreatedAt))
|
||||
truncBody := truncate(pm.Body, 100)
|
||||
|
||||
body := fmt.Sprintf(
|
||||
"**Reminder**: You have a pending message from %s (%s old). Message: \"%s\". Please claim and process it.",
|
||||
pm.FromAgent, age, truncBody,
|
||||
)
|
||||
|
||||
_, err := w.msgService.SendMessage(ctx, "system", pm.ToAgent, body, SendOptions{
|
||||
Subject: fmt.Sprintf("stalemate-reminder:%d", pm.ID),
|
||||
Priority: 7,
|
||||
Metadata: fmt.Sprintf(`{"stalemate_reminder_for":%d}`, pm.ID),
|
||||
})
|
||||
if err != nil {
|
||||
w.logger.Error("send stalemate reminder failed",
|
||||
"message_id", pm.ID,
|
||||
"to_agent", pm.ToAgent,
|
||||
"error", err,
|
||||
)
|
||||
continue
|
||||
}
|
||||
w.logger.Info("sent stalemate reminder",
|
||||
"message_id", pm.ID,
|
||||
"to_agent", pm.ToAgent,
|
||||
"from_agent", pm.FromAgent,
|
||||
"age", age,
|
||||
)
|
||||
count++
|
||||
}
|
||||
return count
|
||||
}
|
||||
|
||||
// escalatePendingMessages escalates pending messages older than EscalateAfter to #approvals.
|
||||
func (w *StalemateWorker) escalatePendingMessages(ctx context.Context) int64 {
|
||||
cutoff := time.Now().Add(-w.config.EscalateAfter)
|
||||
|
||||
rows, err := w.db.QueryContext(ctx,
|
||||
`SELECT id, from_agent, to_agent, body, created_at
|
||||
FROM messages
|
||||
WHERE status = 'pending'
|
||||
AND to_agent IS NOT NULL
|
||||
AND to_agent != ''
|
||||
AND from_agent != 'system'
|
||||
AND to_agent != 'system'
|
||||
AND created_at < ?`,
|
||||
cutoff,
|
||||
)
|
||||
if err != nil {
|
||||
w.logger.Error("query escalation candidates failed", "error", err)
|
||||
return 0
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
type pendingMsg struct {
|
||||
ID int64
|
||||
FromAgent string
|
||||
ToAgent string
|
||||
Body string
|
||||
CreatedAt time.Time
|
||||
}
|
||||
|
||||
var pending []pendingMsg
|
||||
for rows.Next() {
|
||||
var pm pendingMsg
|
||||
if err := rows.Scan(&pm.ID, &pm.FromAgent, &pm.ToAgent, &pm.Body, &pm.CreatedAt); err != nil {
|
||||
w.logger.Error("scan escalation candidate failed", "error", err)
|
||||
continue
|
||||
}
|
||||
pending = append(pending, pm)
|
||||
}
|
||||
|
||||
if len(pending) == 0 {
|
||||
return 0
|
||||
}
|
||||
|
||||
// Look up #approvals channel
|
||||
channelID, err := w.channelLookup.GetChannelIDByName(ctx, "approvals")
|
||||
if err != nil {
|
||||
w.logger.Warn("cannot escalate: #approvals channel not found", "error", err)
|
||||
return 0
|
||||
}
|
||||
|
||||
count := int64(0)
|
||||
for _, pm := range pending {
|
||||
// Check if already escalated
|
||||
if w.escalationExists(ctx, pm.ID) {
|
||||
continue
|
||||
}
|
||||
|
||||
age := formatAge(time.Since(pm.CreatedAt))
|
||||
truncBody := truncate(pm.Body, 100)
|
||||
|
||||
body := fmt.Sprintf(
|
||||
"**ESCALATION**: Pending message for @%s from %s has been unprocessed for %s. Message: \"%s\". Manual intervention may be required.",
|
||||
pm.ToAgent, pm.FromAgent, age, truncBody,
|
||||
)
|
||||
|
||||
_, err := w.msgService.SendMessage(ctx, "system", "", body, SendOptions{
|
||||
Subject: fmt.Sprintf("stalemate-escalation:%d", pm.ID),
|
||||
Priority: 9,
|
||||
Metadata: fmt.Sprintf(`{"stalemate_escalation_for":%d}`, pm.ID),
|
||||
ChannelID: &channelID,
|
||||
})
|
||||
if err != nil {
|
||||
w.logger.Error("send escalation to #approvals failed",
|
||||
"message_id", pm.ID,
|
||||
"error", err,
|
||||
)
|
||||
continue
|
||||
}
|
||||
w.logger.Info("escalated stale message to #approvals",
|
||||
"message_id", pm.ID,
|
||||
"to_agent", pm.ToAgent,
|
||||
"from_agent", pm.FromAgent,
|
||||
"age", age,
|
||||
)
|
||||
count++
|
||||
}
|
||||
return count
|
||||
}
|
||||
|
||||
// reminderExists checks if a system reminder already exists for a given message ID.
|
||||
func (w *StalemateWorker) reminderExists(ctx context.Context, messageID int64, toAgent string) bool {
|
||||
var count int
|
||||
err := w.db.QueryRowContext(ctx,
|
||||
`SELECT COUNT(*) FROM messages
|
||||
WHERE from_agent = 'system'
|
||||
AND to_agent = ?
|
||||
AND metadata LIKE ?`,
|
||||
toAgent, fmt.Sprintf(`%%"stalemate_reminder_for":%d%%`, messageID),
|
||||
).Scan(&count)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return count > 0
|
||||
}
|
||||
|
||||
// escalationExists checks if an escalation already exists for a given message ID.
|
||||
func (w *StalemateWorker) escalationExists(ctx context.Context, messageID int64) bool {
|
||||
var count int
|
||||
err := w.db.QueryRowContext(ctx,
|
||||
`SELECT COUNT(*) FROM messages
|
||||
WHERE from_agent = 'system'
|
||||
AND metadata LIKE ?`,
|
||||
fmt.Sprintf(`%%"stalemate_escalation_for":%d%%`, messageID),
|
||||
).Scan(&count)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return count > 0
|
||||
}
|
||||
|
||||
// truncate truncates a string to maxLen characters, appending "..." if truncated.
|
||||
func truncate(s string, maxLen int) string {
|
||||
runes := []rune(s)
|
||||
if len(runes) <= maxLen {
|
||||
return s
|
||||
}
|
||||
return string(runes[:maxLen]) + "..."
|
||||
}
|
||||
|
||||
// formatAge returns a human-readable age string.
|
||||
func formatAge(d time.Duration) string {
|
||||
if d < time.Hour {
|
||||
return fmt.Sprintf("%dm", int(d.Minutes()))
|
||||
}
|
||||
hours := int(d.Hours())
|
||||
if hours < 24 {
|
||||
return fmt.Sprintf("%dh", hours)
|
||||
}
|
||||
days := hours / 24
|
||||
remainingHours := hours % 24
|
||||
if remainingHours == 0 {
|
||||
if days == 1 {
|
||||
return "1 day"
|
||||
}
|
||||
return fmt.Sprintf("%d days", days)
|
||||
}
|
||||
if days == 1 {
|
||||
return fmt.Sprintf("1 day %dh", remainingHours)
|
||||
}
|
||||
return fmt.Sprintf("%d days %dh", days, remainingHours)
|
||||
}
|
||||
@@ -0,0 +1,480 @@
|
||||
package messaging
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"os"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
|
||||
"github.com/synapbus/synapbus/internal/trace"
|
||||
)
|
||||
|
||||
// stubChannelLookup implements ChannelLookup for tests.
|
||||
type stubChannelLookup struct {
|
||||
channelID int64
|
||||
err error
|
||||
}
|
||||
|
||||
func (s *stubChannelLookup) GetChannelIDByName(ctx context.Context, name string) (int64, error) {
|
||||
if s.err != nil {
|
||||
return 0, s.err
|
||||
}
|
||||
return s.channelID, nil
|
||||
}
|
||||
|
||||
// newStalemateTestService creates a MessagingService and DB for stalemate tests.
|
||||
func newStalemateTestService(t *testing.T) (*MessagingService, *sql.DB) {
|
||||
t.Helper()
|
||||
db := newTestDB(t)
|
||||
|
||||
seedAgent(t, db, "sender")
|
||||
seedAgent(t, db, "receiver")
|
||||
seedAgent(t, db, "system")
|
||||
|
||||
store := NewSQLiteMessageStore(db)
|
||||
tracer := trace.NewTracer(db)
|
||||
t.Cleanup(func() { tracer.Close() })
|
||||
|
||||
svc := NewMessagingService(store, tracer)
|
||||
return svc, db
|
||||
}
|
||||
|
||||
// insertStaleMessage inserts a message with a specific created_at and claimed_at for testing.
|
||||
func insertStaleMessage(t *testing.T, db *sql.DB, from, to, body, status string, createdAt time.Time, claimedAt *time.Time, claimedBy string) int64 {
|
||||
t.Helper()
|
||||
|
||||
// Insert conversation first
|
||||
result, err := db.Exec(
|
||||
`INSERT INTO conversations (subject, created_by, created_at, updated_at)
|
||||
VALUES (?, ?, ?, ?)`,
|
||||
"stalemate-test", from, createdAt, createdAt,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("insert conversation: %v", err)
|
||||
}
|
||||
convID, _ := result.LastInsertId()
|
||||
|
||||
var claimedAtSQL interface{} = nil
|
||||
if claimedAt != nil {
|
||||
claimedAtSQL = *claimedAt
|
||||
}
|
||||
var claimedBySQL interface{} = nil
|
||||
if claimedBy != "" {
|
||||
claimedBySQL = claimedBy
|
||||
}
|
||||
|
||||
result, err = db.Exec(
|
||||
`INSERT INTO messages (conversation_id, from_agent, to_agent, body, priority, status, metadata, claimed_by, claimed_at, created_at, updated_at)
|
||||
VALUES (?, ?, ?, ?, 5, ?, '{}', ?, ?, ?, ?)`,
|
||||
convID, from, to, body, status, claimedBySQL, claimedAtSQL, createdAt, createdAt,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("insert stale message: %v", err)
|
||||
}
|
||||
id, _ := result.LastInsertId()
|
||||
return id
|
||||
}
|
||||
|
||||
func TestStalemateWorker_ProcessingTimeout(t *testing.T) {
|
||||
svc, db := newStalemateTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// Insert a message in "processing" status with old claimed_at
|
||||
oldClaimedAt := time.Now().Add(-25 * time.Hour)
|
||||
msgID := insertStaleMessage(t, db, "sender", "receiver", "stale processing task", StatusProcessing, time.Now().Add(-26*time.Hour), &oldClaimedAt, "receiver")
|
||||
|
||||
config := DefaultStalemateConfig()
|
||||
config.ProcessingTimeout = 24 * time.Hour
|
||||
|
||||
lookup := &stubChannelLookup{channelID: 0, err: fmt.Errorf("no channel")}
|
||||
worker := NewStalemateWorker(db, svc, lookup, config)
|
||||
|
||||
worker.checkStaleMessages(ctx)
|
||||
|
||||
// Verify message was auto-failed
|
||||
var status, metadata string
|
||||
err := db.QueryRowContext(ctx, `SELECT status, metadata FROM messages WHERE id = ?`, msgID).Scan(&status, &metadata)
|
||||
if err != nil {
|
||||
t.Fatalf("query message: %v", err)
|
||||
}
|
||||
if status != StatusFailed {
|
||||
t.Errorf("status = %q, want %q", status, StatusFailed)
|
||||
}
|
||||
if metadata == "{}" {
|
||||
t.Error("expected metadata to contain error info")
|
||||
}
|
||||
}
|
||||
|
||||
func TestStalemateWorker_ProcessingTimeout_NotExpired(t *testing.T) {
|
||||
svc, db := newStalemateTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// Insert a message in "processing" status with recent claimed_at (should NOT be failed)
|
||||
recentClaimedAt := time.Now().Add(-1 * time.Hour)
|
||||
msgID := insertStaleMessage(t, db, "sender", "receiver", "recent processing task", StatusProcessing, time.Now().Add(-2*time.Hour), &recentClaimedAt, "receiver")
|
||||
|
||||
config := DefaultStalemateConfig()
|
||||
config.ProcessingTimeout = 24 * time.Hour
|
||||
|
||||
lookup := &stubChannelLookup{channelID: 0, err: fmt.Errorf("no channel")}
|
||||
worker := NewStalemateWorker(db, svc, lookup, config)
|
||||
|
||||
worker.checkStaleMessages(ctx)
|
||||
|
||||
// Verify message was NOT auto-failed
|
||||
var status string
|
||||
err := db.QueryRowContext(ctx, `SELECT status FROM messages WHERE id = ?`, msgID).Scan(&status)
|
||||
if err != nil {
|
||||
t.Fatalf("query message: %v", err)
|
||||
}
|
||||
if status != StatusProcessing {
|
||||
t.Errorf("status = %q, want %q (should not have been failed)", status, StatusProcessing)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStalemateWorker_PendingReminder(t *testing.T) {
|
||||
svc, db := newStalemateTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// Insert a pending DM that is 5 hours old
|
||||
insertStaleMessage(t, db, "sender", "receiver", "please review this", StatusPending, time.Now().Add(-5*time.Hour), nil, "")
|
||||
|
||||
config := DefaultStalemateConfig()
|
||||
config.ReminderAfter = 4 * time.Hour
|
||||
config.EscalateAfter = 48 * time.Hour // won't trigger
|
||||
|
||||
lookup := &stubChannelLookup{channelID: 0, err: fmt.Errorf("no channel")}
|
||||
worker := NewStalemateWorker(db, svc, lookup, config)
|
||||
|
||||
worker.checkStaleMessages(ctx)
|
||||
|
||||
// Verify a system reminder was sent to receiver
|
||||
var count int
|
||||
err := db.QueryRowContext(ctx,
|
||||
`SELECT COUNT(*) FROM messages WHERE from_agent = 'system' AND to_agent = 'receiver' AND body LIKE '%Reminder%'`,
|
||||
).Scan(&count)
|
||||
if err != nil {
|
||||
t.Fatalf("query reminder: %v", err)
|
||||
}
|
||||
if count != 1 {
|
||||
t.Errorf("expected 1 reminder, got %d", count)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStalemateWorker_SystemMessageSkip(t *testing.T) {
|
||||
svc, db := newStalemateTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// Insert a pending DM FROM system (should be skipped)
|
||||
insertStaleMessage(t, db, "system", "receiver", "system notification", StatusPending, time.Now().Add(-5*time.Hour), nil, "")
|
||||
|
||||
config := DefaultStalemateConfig()
|
||||
config.ReminderAfter = 4 * time.Hour
|
||||
|
||||
lookup := &stubChannelLookup{channelID: 0, err: fmt.Errorf("no channel")}
|
||||
worker := NewStalemateWorker(db, svc, lookup, config)
|
||||
|
||||
worker.checkStaleMessages(ctx)
|
||||
|
||||
// Verify NO reminder was sent (only the original system message should exist)
|
||||
var count int
|
||||
err := db.QueryRowContext(ctx,
|
||||
`SELECT COUNT(*) FROM messages WHERE from_agent = 'system' AND body LIKE '%Reminder%'`,
|
||||
).Scan(&count)
|
||||
if err != nil {
|
||||
t.Fatalf("query reminder: %v", err)
|
||||
}
|
||||
if count != 0 {
|
||||
t.Errorf("expected 0 reminders for system message, got %d", count)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStalemateWorker_DuplicateReminderPrevention(t *testing.T) {
|
||||
svc, db := newStalemateTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// Insert a pending DM that is old enough for a reminder
|
||||
insertStaleMessage(t, db, "sender", "receiver", "need your attention", StatusPending, time.Now().Add(-5*time.Hour), nil, "")
|
||||
|
||||
config := DefaultStalemateConfig()
|
||||
config.ReminderAfter = 4 * time.Hour
|
||||
config.EscalateAfter = 48 * time.Hour
|
||||
|
||||
lookup := &stubChannelLookup{channelID: 0, err: fmt.Errorf("no channel")}
|
||||
worker := NewStalemateWorker(db, svc, lookup, config)
|
||||
|
||||
// Run check twice
|
||||
worker.checkStaleMessages(ctx)
|
||||
worker.checkStaleMessages(ctx)
|
||||
|
||||
// Verify only ONE reminder was sent
|
||||
var count int
|
||||
err := db.QueryRowContext(ctx,
|
||||
`SELECT COUNT(*) FROM messages WHERE from_agent = 'system' AND to_agent = 'receiver' AND body LIKE '%Reminder%'`,
|
||||
).Scan(&count)
|
||||
if err != nil {
|
||||
t.Fatalf("query reminders: %v", err)
|
||||
}
|
||||
if count != 1 {
|
||||
t.Errorf("expected 1 reminder (no duplicates), got %d", count)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStalemateWorker_Escalation(t *testing.T) {
|
||||
svc, db := newStalemateTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// Create #approvals channel
|
||||
_, err := db.Exec(
|
||||
`INSERT INTO channels (id, name, description, topic, type, is_private, is_system, created_by, created_at, updated_at)
|
||||
VALUES (1, 'approvals', 'Approval queue', '', 'standard', 0, 0, 'system', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`)
|
||||
if err != nil {
|
||||
t.Fatalf("create approvals channel: %v", err)
|
||||
}
|
||||
// Add system as member
|
||||
_, err = db.Exec(
|
||||
`INSERT INTO channel_members (channel_id, agent_name, role, joined_at)
|
||||
VALUES (1, 'system', 'owner', CURRENT_TIMESTAMP)`)
|
||||
if err != nil {
|
||||
t.Fatalf("add system to channel: %v", err)
|
||||
}
|
||||
|
||||
// Insert a pending DM that is 49 hours old (beyond escalation threshold)
|
||||
insertStaleMessage(t, db, "sender", "receiver", "urgent task ignored", StatusPending, time.Now().Add(-49*time.Hour), nil, "")
|
||||
|
||||
config := DefaultStalemateConfig()
|
||||
config.ReminderAfter = 4 * time.Hour
|
||||
config.EscalateAfter = 48 * time.Hour
|
||||
|
||||
lookup := &stubChannelLookup{channelID: 1}
|
||||
worker := NewStalemateWorker(db, svc, lookup, config)
|
||||
|
||||
worker.checkStaleMessages(ctx)
|
||||
|
||||
// Verify an escalation was sent to #approvals channel
|
||||
var count int
|
||||
err = db.QueryRowContext(ctx,
|
||||
`SELECT COUNT(*) FROM messages WHERE from_agent = 'system' AND channel_id = 1 AND body LIKE '%ESCALATION%'`,
|
||||
).Scan(&count)
|
||||
if err != nil {
|
||||
t.Fatalf("query escalations: %v", err)
|
||||
}
|
||||
if count != 1 {
|
||||
t.Errorf("expected 1 escalation, got %d", count)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStalemateWorker_DuplicateEscalationPrevention(t *testing.T) {
|
||||
svc, db := newStalemateTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// Create #approvals channel
|
||||
db.Exec(
|
||||
`INSERT INTO channels (id, name, description, topic, type, is_private, is_system, created_by, created_at, updated_at)
|
||||
VALUES (1, 'approvals', 'Approval queue', '', 'standard', 0, 0, 'system', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`)
|
||||
db.Exec(
|
||||
`INSERT INTO channel_members (channel_id, agent_name, role, joined_at)
|
||||
VALUES (1, 'system', 'owner', CURRENT_TIMESTAMP)`)
|
||||
|
||||
// Insert a pending DM that is 49 hours old
|
||||
insertStaleMessage(t, db, "sender", "receiver", "urgent task", StatusPending, time.Now().Add(-49*time.Hour), nil, "")
|
||||
|
||||
config := DefaultStalemateConfig()
|
||||
config.ReminderAfter = 4 * time.Hour
|
||||
config.EscalateAfter = 48 * time.Hour
|
||||
|
||||
lookup := &stubChannelLookup{channelID: 1}
|
||||
worker := NewStalemateWorker(db, svc, lookup, config)
|
||||
|
||||
// Run check twice
|
||||
worker.checkStaleMessages(ctx)
|
||||
worker.checkStaleMessages(ctx)
|
||||
|
||||
// Verify only ONE escalation was sent
|
||||
var count int
|
||||
err := db.QueryRowContext(ctx,
|
||||
`SELECT COUNT(*) FROM messages WHERE from_agent = 'system' AND channel_id = 1 AND body LIKE '%ESCALATION%'`,
|
||||
).Scan(&count)
|
||||
if err != nil {
|
||||
t.Fatalf("query escalations: %v", err)
|
||||
}
|
||||
if count != 1 {
|
||||
t.Errorf("expected 1 escalation (no duplicates), got %d", count)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseStalemateConfig(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
envVars map[string]string
|
||||
expected StalemateConfig
|
||||
}{
|
||||
{
|
||||
name: "defaults when no env vars",
|
||||
envVars: map[string]string{},
|
||||
expected: DefaultStalemateConfig(),
|
||||
},
|
||||
{
|
||||
name: "custom values with day format",
|
||||
envVars: map[string]string{
|
||||
"SYNAPBUS_STALEMATE_PROCESSING_TIMEOUT": "7d",
|
||||
"SYNAPBUS_STALEMATE_REMINDER_AFTER": "8h",
|
||||
"SYNAPBUS_STALEMATE_ESCALATE_AFTER": "3d",
|
||||
"SYNAPBUS_STALEMATE_INTERVAL": "30m",
|
||||
},
|
||||
expected: StalemateConfig{
|
||||
ProcessingTimeout: 7 * 24 * time.Hour,
|
||||
ReminderAfter: 8 * time.Hour,
|
||||
EscalateAfter: 3 * 24 * time.Hour,
|
||||
Interval: 30 * time.Minute,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "standard Go duration format",
|
||||
envVars: map[string]string{
|
||||
"SYNAPBUS_STALEMATE_PROCESSING_TIMEOUT": "48h",
|
||||
"SYNAPBUS_STALEMATE_REMINDER_AFTER": "2h30m",
|
||||
"SYNAPBUS_STALEMATE_ESCALATE_AFTER": "72h",
|
||||
"SYNAPBUS_STALEMATE_INTERVAL": "5m",
|
||||
},
|
||||
expected: StalemateConfig{
|
||||
ProcessingTimeout: 48 * time.Hour,
|
||||
ReminderAfter: 2*time.Hour + 30*time.Minute,
|
||||
EscalateAfter: 72 * time.Hour,
|
||||
Interval: 5 * time.Minute,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "invalid values fall back to defaults",
|
||||
envVars: map[string]string{
|
||||
"SYNAPBUS_STALEMATE_PROCESSING_TIMEOUT": "invalid",
|
||||
"SYNAPBUS_STALEMATE_REMINDER_AFTER": "bad",
|
||||
"SYNAPBUS_STALEMATE_ESCALATE_AFTER": "",
|
||||
"SYNAPBUS_STALEMATE_INTERVAL": "-5m",
|
||||
},
|
||||
expected: DefaultStalemateConfig(),
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
// Clear all env vars first
|
||||
envKeys := []string{
|
||||
"SYNAPBUS_STALEMATE_PROCESSING_TIMEOUT",
|
||||
"SYNAPBUS_STALEMATE_REMINDER_AFTER",
|
||||
"SYNAPBUS_STALEMATE_ESCALATE_AFTER",
|
||||
"SYNAPBUS_STALEMATE_INTERVAL",
|
||||
}
|
||||
for _, k := range envKeys {
|
||||
os.Unsetenv(k)
|
||||
}
|
||||
|
||||
// Set test env vars
|
||||
for k, v := range tt.envVars {
|
||||
os.Setenv(k, v)
|
||||
}
|
||||
defer func() {
|
||||
for _, k := range envKeys {
|
||||
os.Unsetenv(k)
|
||||
}
|
||||
}()
|
||||
|
||||
cfg := ParseStalemateConfig()
|
||||
|
||||
if cfg.ProcessingTimeout != tt.expected.ProcessingTimeout {
|
||||
t.Errorf("ProcessingTimeout = %v, want %v", cfg.ProcessingTimeout, tt.expected.ProcessingTimeout)
|
||||
}
|
||||
if cfg.ReminderAfter != tt.expected.ReminderAfter {
|
||||
t.Errorf("ReminderAfter = %v, want %v", cfg.ReminderAfter, tt.expected.ReminderAfter)
|
||||
}
|
||||
if cfg.EscalateAfter != tt.expected.EscalateAfter {
|
||||
t.Errorf("EscalateAfter = %v, want %v", cfg.EscalateAfter, tt.expected.EscalateAfter)
|
||||
}
|
||||
if cfg.Interval != tt.expected.Interval {
|
||||
t.Errorf("Interval = %v, want %v", cfg.Interval, tt.expected.Interval)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseDurationWithDays(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
input string
|
||||
want time.Duration
|
||||
wantErr bool
|
||||
}{
|
||||
{"7 days", "7d", 7 * 24 * time.Hour, false},
|
||||
{"1 day", "1d", 24 * time.Hour, false},
|
||||
{"30 days", "30d", 30 * 24 * time.Hour, false},
|
||||
{"standard hours", "48h", 48 * time.Hour, false},
|
||||
{"standard minutes", "15m", 15 * time.Minute, false},
|
||||
{"mixed duration", "2h30m", 2*time.Hour + 30*time.Minute, false},
|
||||
{"empty string", "", 0, true},
|
||||
{"invalid", "xyz", 0, true},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got, err := parseDurationWithDays(tt.input)
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Errorf("parseDurationWithDays(%q) error = %v, wantErr %v", tt.input, err, tt.wantErr)
|
||||
return
|
||||
}
|
||||
if got != tt.want {
|
||||
t.Errorf("parseDurationWithDays(%q) = %v, want %v", tt.input, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestTruncate(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
input string
|
||||
maxLen int
|
||||
want string
|
||||
}{
|
||||
{"short string", "hello", 10, "hello"},
|
||||
{"exact length", "hello", 5, "hello"},
|
||||
{"truncated", "hello world, this is a long message", 10, "hello worl..."},
|
||||
{"empty", "", 10, ""},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got := truncate(tt.input, tt.maxLen)
|
||||
if got != tt.want {
|
||||
t.Errorf("truncate(%q, %d) = %q, want %q", tt.input, tt.maxLen, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestFormatAge(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
d time.Duration
|
||||
want string
|
||||
}{
|
||||
{"minutes", 30 * time.Minute, "30m"},
|
||||
{"hours", 5 * time.Hour, "5h"},
|
||||
{"1 day", 24 * time.Hour, "1 day"},
|
||||
{"2 days", 48 * time.Hour, "2 days"},
|
||||
{"1 day with hours", 25 * time.Hour, "1 day 1h"},
|
||||
{"2 days with hours", 50 * time.Hour, "2 days 2h"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got := formatAge(tt.d)
|
||||
if got != tt.want {
|
||||
t.Errorf("formatAge(%v) = %q, want %q", tt.d, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -39,6 +39,7 @@ type MessageStore interface {
|
||||
GetLastReadForDM(ctx context.Context, agentNames []string, peerAgent string) (int64, error)
|
||||
GetConversationIDsForChannel(ctx context.Context, channelID int64, lastMessageID int64) ([]int64, error)
|
||||
GetConversationIDsForDM(ctx context.Context, agentNames []string, peerAgent string, lastMessageID int64) ([]int64, error)
|
||||
GetReplyCounts(ctx context.Context, messageIDs []int64) (map[int64]int, error)
|
||||
}
|
||||
|
||||
// SQLiteMessageStore implements MessageStore using SQLite.
|
||||
@@ -492,6 +493,57 @@ func (s *SQLiteMessageStore) buildSearchConditions(agentName, query string, opts
|
||||
args = append(args, opts.Channel)
|
||||
}
|
||||
|
||||
if len(opts.Channels) > 0 {
|
||||
placeholders := make([]string, len(opts.Channels))
|
||||
for i, ch := range opts.Channels {
|
||||
placeholders[i] = "?"
|
||||
args = append(args, ch)
|
||||
}
|
||||
conditions = append(conditions, fmt.Sprintf("m.channel_id IN (SELECT id FROM channels WHERE LOWER(name) IN (%s))", strings.Join(placeholders, ",")))
|
||||
}
|
||||
|
||||
if len(opts.ExcludeChannels) > 0 {
|
||||
placeholders := make([]string, len(opts.ExcludeChannels))
|
||||
for i, ch := range opts.ExcludeChannels {
|
||||
placeholders[i] = "?"
|
||||
args = append(args, ch)
|
||||
}
|
||||
conditions = append(conditions, fmt.Sprintf("(m.channel_id IS NULL OR m.channel_id NOT IN (SELECT id FROM channels WHERE LOWER(name) IN (%s)))", strings.Join(placeholders, ",")))
|
||||
}
|
||||
|
||||
if len(opts.Agents) > 0 {
|
||||
placeholders := make([]string, len(opts.Agents))
|
||||
for i, a := range opts.Agents {
|
||||
placeholders[i] = "?"
|
||||
args = append(args, a)
|
||||
}
|
||||
inClause := strings.Join(placeholders, ",")
|
||||
// Clone placeholders for the second IN clause
|
||||
placeholders2 := make([]string, len(opts.Agents))
|
||||
for i, a := range opts.Agents {
|
||||
placeholders2[i] = "?"
|
||||
args = append(args, a)
|
||||
}
|
||||
inClause2 := strings.Join(placeholders2, ",")
|
||||
conditions = append(conditions, fmt.Sprintf("(m.from_agent IN (%s) OR m.to_agent IN (%s))", inClause, inClause2))
|
||||
}
|
||||
|
||||
if len(opts.ExcludeAgents) > 0 {
|
||||
placeholders := make([]string, len(opts.ExcludeAgents))
|
||||
for i, a := range opts.ExcludeAgents {
|
||||
placeholders[i] = "?"
|
||||
args = append(args, a)
|
||||
}
|
||||
inClause := strings.Join(placeholders, ",")
|
||||
placeholders2 := make([]string, len(opts.ExcludeAgents))
|
||||
for i, a := range opts.ExcludeAgents {
|
||||
placeholders2[i] = "?"
|
||||
args = append(args, a)
|
||||
}
|
||||
inClause2 := strings.Join(placeholders2, ",")
|
||||
conditions = append(conditions, fmt.Sprintf("m.from_agent NOT IN (%s) AND (m.to_agent = '' OR m.to_agent NOT IN (%s))", inClause, inClause2))
|
||||
}
|
||||
|
||||
if opts.After != "" {
|
||||
if t, err := time.Parse(time.RFC3339, opts.After); err == nil {
|
||||
conditions = append(conditions, "m.created_at >= ?")
|
||||
@@ -956,6 +1008,44 @@ func (s *SQLiteMessageStore) GetConversationIDsForDM(ctx context.Context, agentN
|
||||
return ids, rows.Err()
|
||||
}
|
||||
|
||||
// GetReplyCounts returns a map of message ID → reply count for the given IDs.
|
||||
func (s *SQLiteMessageStore) GetReplyCounts(ctx context.Context, messageIDs []int64) (map[int64]int, error) {
|
||||
if len(messageIDs) == 0 {
|
||||
return map[int64]int{}, nil
|
||||
}
|
||||
|
||||
placeholders := make([]string, len(messageIDs))
|
||||
args := make([]any, len(messageIDs))
|
||||
for i, id := range messageIDs {
|
||||
placeholders[i] = "?"
|
||||
args[i] = id
|
||||
}
|
||||
|
||||
query := fmt.Sprintf(
|
||||
`SELECT reply_to, COUNT(*) FROM messages
|
||||
WHERE reply_to IN (%s)
|
||||
GROUP BY reply_to`,
|
||||
strings.Join(placeholders, ","),
|
||||
)
|
||||
|
||||
rows, err := s.db.QueryContext(ctx, query, args...)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get reply counts: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
counts := make(map[int64]int)
|
||||
for rows.Next() {
|
||||
var replyTo int64
|
||||
var count int
|
||||
if err := rows.Scan(&replyTo, &count); err != nil {
|
||||
return nil, fmt.Errorf("scan reply count: %w", err)
|
||||
}
|
||||
counts[replyTo] = count
|
||||
}
|
||||
return counts, rows.Err()
|
||||
}
|
||||
|
||||
// scanMessage scans a single message from sql.Row.
|
||||
func scanMessage(row *sql.Row) (*Message, error) {
|
||||
var msg Message
|
||||
|
||||
@@ -1027,6 +1027,132 @@ func TestSQLiteMessageStore_GetChannelMessages_Offset(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestSQLiteMessageStore_GetReplyCounts(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteMessageStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
seedAgent(t, db, "sender")
|
||||
seedAgent(t, db, "replier")
|
||||
|
||||
conv := &Conversation{Subject: "reply counts", CreatedBy: "sender"}
|
||||
if err := store.InsertConversation(ctx, conv); err != nil {
|
||||
t.Fatalf("InsertConversation: %v", err)
|
||||
}
|
||||
|
||||
// Insert parent message
|
||||
parent := &Message{
|
||||
ConversationID: conv.ID,
|
||||
FromAgent: "sender",
|
||||
ToAgent: "replier",
|
||||
Body: "parent message",
|
||||
Priority: 5,
|
||||
Status: StatusPending,
|
||||
}
|
||||
if err := store.InsertMessage(ctx, parent); err != nil {
|
||||
t.Fatalf("InsertMessage (parent): %v", err)
|
||||
}
|
||||
|
||||
// Insert 3 replies to parent
|
||||
for i := 0; i < 3; i++ {
|
||||
reply := &Message{
|
||||
ConversationID: conv.ID,
|
||||
FromAgent: "replier",
|
||||
ToAgent: "sender",
|
||||
ReplyTo: &parent.ID,
|
||||
Body: fmt.Sprintf("reply %d", i),
|
||||
Priority: 5,
|
||||
Status: StatusPending,
|
||||
}
|
||||
if err := store.InsertMessage(ctx, reply); err != nil {
|
||||
t.Fatalf("InsertMessage (reply %d): %v", i, err)
|
||||
}
|
||||
}
|
||||
|
||||
t.Run("parent with 3 replies", func(t *testing.T) {
|
||||
counts, err := store.GetReplyCounts(ctx, []int64{parent.ID})
|
||||
if err != nil {
|
||||
t.Fatalf("GetReplyCounts: %v", err)
|
||||
}
|
||||
if counts[parent.ID] != 3 {
|
||||
t.Errorf("reply count for parent = %d, want 3", counts[parent.ID])
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("message with no replies returns 0", func(t *testing.T) {
|
||||
// Insert a message with no replies
|
||||
noReply := &Message{
|
||||
ConversationID: conv.ID,
|
||||
FromAgent: "sender",
|
||||
ToAgent: "replier",
|
||||
Body: "no replies here",
|
||||
Priority: 5,
|
||||
Status: StatusPending,
|
||||
}
|
||||
if err := store.InsertMessage(ctx, noReply); err != nil {
|
||||
t.Fatalf("InsertMessage (noReply): %v", err)
|
||||
}
|
||||
|
||||
counts, err := store.GetReplyCounts(ctx, []int64{noReply.ID})
|
||||
if err != nil {
|
||||
t.Fatalf("GetReplyCounts: %v", err)
|
||||
}
|
||||
if counts[noReply.ID] != 0 {
|
||||
t.Errorf("reply count for noReply = %d, want 0", counts[noReply.ID])
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("multiple parents", func(t *testing.T) {
|
||||
// Insert a second parent with 2 replies
|
||||
parent2 := &Message{
|
||||
ConversationID: conv.ID,
|
||||
FromAgent: "sender",
|
||||
ToAgent: "replier",
|
||||
Body: "second parent",
|
||||
Priority: 5,
|
||||
Status: StatusPending,
|
||||
}
|
||||
if err := store.InsertMessage(ctx, parent2); err != nil {
|
||||
t.Fatalf("InsertMessage (parent2): %v", err)
|
||||
}
|
||||
for i := 0; i < 2; i++ {
|
||||
reply := &Message{
|
||||
ConversationID: conv.ID,
|
||||
FromAgent: "replier",
|
||||
ToAgent: "sender",
|
||||
ReplyTo: &parent2.ID,
|
||||
Body: fmt.Sprintf("reply to parent2 %d", i),
|
||||
Priority: 5,
|
||||
Status: StatusPending,
|
||||
}
|
||||
if err := store.InsertMessage(ctx, reply); err != nil {
|
||||
t.Fatalf("InsertMessage (parent2 reply %d): %v", i, err)
|
||||
}
|
||||
}
|
||||
|
||||
counts, err := store.GetReplyCounts(ctx, []int64{parent.ID, parent2.ID})
|
||||
if err != nil {
|
||||
t.Fatalf("GetReplyCounts: %v", err)
|
||||
}
|
||||
if counts[parent.ID] != 3 {
|
||||
t.Errorf("reply count for parent = %d, want 3", counts[parent.ID])
|
||||
}
|
||||
if counts[parent2.ID] != 2 {
|
||||
t.Errorf("reply count for parent2 = %d, want 2", counts[parent2.ID])
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("empty slice returns empty map", func(t *testing.T) {
|
||||
counts, err := store.GetReplyCounts(ctx, []int64{})
|
||||
if err != nil {
|
||||
t.Fatalf("GetReplyCounts: %v", err)
|
||||
}
|
||||
if len(counts) != 0 {
|
||||
t.Errorf("expected empty map, got %v", counts)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestSQLiteMessageStore_CombinedFiltersAndPagination(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteMessageStore(db)
|
||||
|
||||
@@ -14,6 +14,16 @@ const (
|
||||
StatusFailed = "failed"
|
||||
)
|
||||
|
||||
// AttachmentInfo is a lightweight attachment summary included in message responses.
|
||||
// It avoids importing the attachments package into the messaging package.
|
||||
type AttachmentInfo struct {
|
||||
Hash string `json:"hash"`
|
||||
OriginalFilename string `json:"original_filename"`
|
||||
Size int64 `json:"size"`
|
||||
MIMEType string `json:"mime_type"`
|
||||
IsImage bool `json:"is_image"`
|
||||
}
|
||||
|
||||
// Message represents a single message in the system.
|
||||
type Message struct {
|
||||
ID int64 `json:"id"`
|
||||
@@ -30,6 +40,18 @@ type Message struct {
|
||||
ClaimedAt *time.Time `json:"claimed_at,omitempty"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
ReplyCount int `json:"reply_count"`
|
||||
Attachments []AttachmentInfo `json:"attachments,omitempty"`
|
||||
WorkflowState string `json:"workflow_state,omitempty"`
|
||||
Reactions []ReactionInfo `json:"reactions,omitempty"`
|
||||
}
|
||||
|
||||
// ReactionInfo is a lightweight reaction summary included in message responses.
|
||||
type ReactionInfo struct {
|
||||
AgentName string `json:"agent_name"`
|
||||
Reaction string `json:"reaction"`
|
||||
Metadata json.RawMessage `json:"metadata,omitempty"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
|
||||
// Conversation groups related messages into a thread.
|
||||
|
||||
@@ -0,0 +1,279 @@
|
||||
package push
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/ecdsa"
|
||||
"crypto/elliptic"
|
||||
"crypto/rand"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"math/big"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
|
||||
webpush "github.com/SherClockHolmes/webpush-go"
|
||||
)
|
||||
|
||||
// Service manages Web Push notifications and VAPID key lifecycle.
|
||||
type Service struct {
|
||||
store Store
|
||||
vapidKeys *VAPIDKeys
|
||||
dataDir string
|
||||
logger *slog.Logger
|
||||
mu sync.RWMutex
|
||||
}
|
||||
|
||||
// VAPIDKeys holds the VAPID key pair for Web Push authentication.
|
||||
type VAPIDKeys struct {
|
||||
PublicKey string `json:"public_key"`
|
||||
PrivateKey string `json:"private_key"`
|
||||
}
|
||||
|
||||
// Notification represents a push notification payload.
|
||||
type Notification struct {
|
||||
Title string `json:"title"`
|
||||
Body string `json:"body"`
|
||||
Tag string `json:"tag,omitempty"`
|
||||
URL string `json:"url,omitempty"`
|
||||
}
|
||||
|
||||
// NewService creates a new push notification service.
|
||||
// It loads or generates VAPID keys from the data directory.
|
||||
func NewService(store Store, dataDir string, logger *slog.Logger) (*Service, error) {
|
||||
s := &Service{
|
||||
store: store,
|
||||
dataDir: dataDir,
|
||||
logger: logger.With("component", "push"),
|
||||
}
|
||||
|
||||
keys, err := s.loadOrGenerateVAPIDKeys()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("load or generate VAPID keys: %w", err)
|
||||
}
|
||||
s.vapidKeys = keys
|
||||
|
||||
s.logger.Info("push notification service initialized",
|
||||
"vapid_public_key", keys.PublicKey[:16]+"...",
|
||||
)
|
||||
|
||||
return s, nil
|
||||
}
|
||||
|
||||
// GetVAPIDPublicKey returns the VAPID public key for client subscription.
|
||||
func (s *Service) GetVAPIDPublicKey() string {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
return s.vapidKeys.PublicKey
|
||||
}
|
||||
|
||||
// Subscribe registers a push subscription for a user.
|
||||
func (s *Service) Subscribe(ctx context.Context, userID int64, endpoint, keyP256dh, keyAuth, userAgent string) error {
|
||||
return s.store.Subscribe(ctx, userID, endpoint, keyP256dh, keyAuth, userAgent)
|
||||
}
|
||||
|
||||
// Unsubscribe removes a push subscription by endpoint, scoped to user.
|
||||
func (s *Service) Unsubscribe(ctx context.Context, userID int64, endpoint string) error {
|
||||
return s.store.Unsubscribe(ctx, userID, endpoint)
|
||||
}
|
||||
|
||||
// SendToUser sends a push notification to all subscriptions for a user.
|
||||
// It automatically removes subscriptions that return 410 Gone (unsubscribed).
|
||||
func (s *Service) SendToUser(ctx context.Context, userID int64, notification Notification) error {
|
||||
subs, err := s.store.GetSubscriptions(ctx, userID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("get subscriptions: %w", err)
|
||||
}
|
||||
|
||||
if len(subs) == 0 {
|
||||
s.logger.Debug("no push subscriptions for user", "user_id", userID)
|
||||
return nil
|
||||
}
|
||||
|
||||
payload, err := json.Marshal(notification)
|
||||
if err != nil {
|
||||
return fmt.Errorf("marshal notification: %w", err)
|
||||
}
|
||||
|
||||
s.mu.RLock()
|
||||
vapidPrivate := s.vapidKeys.PrivateKey
|
||||
vapidPublic := s.vapidKeys.PublicKey
|
||||
s.mu.RUnlock()
|
||||
|
||||
var sendErrors []error
|
||||
for _, sub := range subs {
|
||||
wpSub := &webpush.Subscription{
|
||||
Endpoint: sub.Endpoint,
|
||||
Keys: webpush.Keys{
|
||||
P256dh: sub.KeyP256dh,
|
||||
Auth: sub.KeyAuth,
|
||||
},
|
||||
}
|
||||
|
||||
resp, err := webpush.SendNotification(payload, wpSub, &webpush.Options{
|
||||
VAPIDPrivateKey: vapidPrivate,
|
||||
VAPIDPublicKey: vapidPublic,
|
||||
Subscriber: "mailto:noreply@synapbus.local",
|
||||
TTL: 86400, // 24 hours
|
||||
})
|
||||
if err != nil {
|
||||
s.logger.Warn("push notification send failed",
|
||||
"user_id", userID,
|
||||
"endpoint", truncateEndpoint(sub.Endpoint),
|
||||
"error", err,
|
||||
)
|
||||
sendErrors = append(sendErrors, err)
|
||||
continue
|
||||
}
|
||||
resp.Body.Close()
|
||||
|
||||
// Remove stale subscriptions
|
||||
if resp.StatusCode == http.StatusGone {
|
||||
s.logger.Info("removing stale push subscription",
|
||||
"user_id", userID,
|
||||
"endpoint", truncateEndpoint(sub.Endpoint),
|
||||
)
|
||||
if delErr := s.store.DeleteSubscription(ctx, sub.Endpoint); delErr != nil {
|
||||
s.logger.Warn("failed to delete stale subscription",
|
||||
"endpoint", truncateEndpoint(sub.Endpoint),
|
||||
"error", delErr,
|
||||
)
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
if resp.StatusCode >= 400 {
|
||||
s.logger.Warn("push notification rejected",
|
||||
"user_id", userID,
|
||||
"endpoint", truncateEndpoint(sub.Endpoint),
|
||||
"status", resp.StatusCode,
|
||||
)
|
||||
sendErrors = append(sendErrors, fmt.Errorf("push endpoint returned %d", resp.StatusCode))
|
||||
|
||||
// Also remove subscriptions that return 404 (endpoint no longer valid)
|
||||
if resp.StatusCode == http.StatusNotFound {
|
||||
if delErr := s.store.DeleteSubscription(ctx, sub.Endpoint); delErr != nil {
|
||||
s.logger.Warn("failed to delete invalid subscription",
|
||||
"endpoint", truncateEndpoint(sub.Endpoint),
|
||||
"error", delErr,
|
||||
)
|
||||
}
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
s.logger.Debug("push notification sent",
|
||||
"user_id", userID,
|
||||
"endpoint", truncateEndpoint(sub.Endpoint),
|
||||
"status", resp.StatusCode,
|
||||
)
|
||||
}
|
||||
|
||||
if len(sendErrors) > 0 {
|
||||
return fmt.Errorf("failed to send to %d/%d subscriptions", len(sendErrors), len(subs))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// vapidKeysPath returns the file path for the VAPID keys JSON file.
|
||||
func (s *Service) vapidKeysPath() string {
|
||||
return filepath.Join(s.dataDir, "vapid_keys.json")
|
||||
}
|
||||
|
||||
// loadOrGenerateVAPIDKeys loads existing VAPID keys from disk or generates new ones.
|
||||
func (s *Service) loadOrGenerateVAPIDKeys() (*VAPIDKeys, error) {
|
||||
keysPath := s.vapidKeysPath()
|
||||
|
||||
// Try to load existing keys
|
||||
data, err := os.ReadFile(keysPath)
|
||||
if err == nil {
|
||||
var keys VAPIDKeys
|
||||
if err := json.Unmarshal(data, &keys); err == nil && keys.PublicKey != "" && keys.PrivateKey != "" {
|
||||
s.logger.Info("loaded existing VAPID keys", "path", keysPath)
|
||||
return &keys, nil
|
||||
}
|
||||
s.logger.Warn("invalid VAPID keys file, regenerating", "path", keysPath)
|
||||
}
|
||||
|
||||
// Generate new VAPID keys
|
||||
keys, err := generateVAPIDKeys()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("generate VAPID keys: %w", err)
|
||||
}
|
||||
|
||||
// Ensure data directory exists
|
||||
if err := os.MkdirAll(s.dataDir, 0o755); err != nil {
|
||||
return nil, fmt.Errorf("create data directory: %w", err)
|
||||
}
|
||||
|
||||
// Save to disk
|
||||
data, err = json.MarshalIndent(keys, "", " ")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("marshal VAPID keys: %w", err)
|
||||
}
|
||||
|
||||
if err := os.WriteFile(keysPath, data, 0o600); err != nil {
|
||||
return nil, fmt.Errorf("write VAPID keys: %w", err)
|
||||
}
|
||||
|
||||
s.logger.Info("generated new VAPID keys", "path", keysPath)
|
||||
return keys, nil
|
||||
}
|
||||
|
||||
// generateVAPIDKeys creates a new ECDSA P-256 key pair and encodes them
|
||||
// as uncompressed (public) and raw (private) base64url strings for VAPID.
|
||||
func generateVAPIDKeys() (*VAPIDKeys, error) {
|
||||
privateKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("generate ECDSA key: %w", err)
|
||||
}
|
||||
|
||||
// Encode public key as uncompressed point (0x04 || x || y)
|
||||
pubBytes := elliptic.Marshal(elliptic.P256(), privateKey.PublicKey.X, privateKey.PublicKey.Y)
|
||||
publicKeyB64 := base64.RawURLEncoding.EncodeToString(pubBytes)
|
||||
|
||||
// Encode private key as raw big-endian bytes (32 bytes, zero-padded)
|
||||
privBytes := privateKey.D.Bytes()
|
||||
// Pad to 32 bytes if needed
|
||||
padded := make([]byte, 32)
|
||||
copy(padded[32-len(privBytes):], privBytes)
|
||||
privateKeyB64 := base64.RawURLEncoding.EncodeToString(padded)
|
||||
|
||||
return &VAPIDKeys{
|
||||
PublicKey: publicKeyB64,
|
||||
PrivateKey: privateKeyB64,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// ParseVAPIDPrivateKey decodes a base64url-encoded VAPID private key
|
||||
// into an *ecdsa.PrivateKey. Useful for testing.
|
||||
func ParseVAPIDPrivateKey(b64 string) (*ecdsa.PrivateKey, error) {
|
||||
raw, err := base64.RawURLEncoding.DecodeString(b64)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("decode private key: %w", err)
|
||||
}
|
||||
|
||||
d := new(big.Int).SetBytes(raw)
|
||||
curve := elliptic.P256()
|
||||
x, y := curve.ScalarBaseMult(raw)
|
||||
|
||||
return &ecdsa.PrivateKey{
|
||||
PublicKey: ecdsa.PublicKey{
|
||||
Curve: curve,
|
||||
X: x,
|
||||
Y: y,
|
||||
},
|
||||
D: d,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// truncateEndpoint returns a truncated version of the endpoint URL for logging.
|
||||
func truncateEndpoint(endpoint string) string {
|
||||
if len(endpoint) <= 40 {
|
||||
return endpoint
|
||||
}
|
||||
return endpoint[:37] + "..."
|
||||
}
|
||||
@@ -0,0 +1,340 @@
|
||||
package push
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/ecdsa"
|
||||
"crypto/elliptic"
|
||||
"crypto/rand"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// mockStore implements Store for testing without a database.
|
||||
type mockStore struct {
|
||||
subscriptions map[int64][]Subscription
|
||||
deleted []string
|
||||
}
|
||||
|
||||
func newMockStore() *mockStore {
|
||||
return &mockStore{
|
||||
subscriptions: make(map[int64][]Subscription),
|
||||
}
|
||||
}
|
||||
|
||||
func (m *mockStore) Subscribe(_ context.Context, userID int64, endpoint, keyP256dh, keyAuth, userAgent string) error {
|
||||
// Upsert: remove existing endpoint first
|
||||
subs := m.subscriptions[userID]
|
||||
for i, s := range subs {
|
||||
if s.Endpoint == endpoint {
|
||||
subs = append(subs[:i], subs[i+1:]...)
|
||||
break
|
||||
}
|
||||
}
|
||||
m.subscriptions[userID] = append(subs, Subscription{
|
||||
UserID: userID,
|
||||
Endpoint: endpoint,
|
||||
KeyP256dh: keyP256dh,
|
||||
KeyAuth: keyAuth,
|
||||
UserAgent: userAgent,
|
||||
})
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *mockStore) Unsubscribe(_ context.Context, userID int64, endpoint string) error {
|
||||
subs := m.subscriptions[userID]
|
||||
for i, s := range subs {
|
||||
if s.Endpoint == endpoint {
|
||||
m.subscriptions[userID] = append(subs[:i], subs[i+1:]...)
|
||||
return nil
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *mockStore) GetSubscriptions(_ context.Context, userID int64) ([]Subscription, error) {
|
||||
return m.subscriptions[userID], nil
|
||||
}
|
||||
|
||||
func (m *mockStore) DeleteSubscription(_ context.Context, endpoint string) error {
|
||||
m.deleted = append(m.deleted, endpoint)
|
||||
// Find and remove from any user's subscriptions
|
||||
for uid, subs := range m.subscriptions {
|
||||
for i, s := range subs {
|
||||
if s.Endpoint == endpoint {
|
||||
m.subscriptions[uid] = append(subs[:i], subs[i+1:]...)
|
||||
return nil
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestVAPIDKeyGeneration(t *testing.T) {
|
||||
keys, err := generateVAPIDKeys()
|
||||
if err != nil {
|
||||
t.Fatalf("generateVAPIDKeys: %v", err)
|
||||
}
|
||||
|
||||
if keys.PublicKey == "" {
|
||||
t.Error("public key is empty")
|
||||
}
|
||||
if keys.PrivateKey == "" {
|
||||
t.Error("private key is empty")
|
||||
}
|
||||
|
||||
// Public key should be 65 bytes (uncompressed point) base64url encoded
|
||||
// 65 bytes -> ceil(65*4/3) = 87 chars without padding
|
||||
if len(keys.PublicKey) < 80 {
|
||||
t.Errorf("public key too short: %d chars", len(keys.PublicKey))
|
||||
}
|
||||
|
||||
// Private key should be 32 bytes base64url encoded
|
||||
// 32 bytes -> ceil(32*4/3) = 43 chars without padding
|
||||
if len(keys.PrivateKey) < 40 {
|
||||
t.Errorf("private key too short: %d chars", len(keys.PrivateKey))
|
||||
}
|
||||
|
||||
// Should be parseable back into an ECDSA key
|
||||
privKey, err := ParseVAPIDPrivateKey(keys.PrivateKey)
|
||||
if err != nil {
|
||||
t.Fatalf("ParseVAPIDPrivateKey: %v", err)
|
||||
}
|
||||
if privKey.Curve == nil {
|
||||
t.Error("parsed key has nil curve")
|
||||
}
|
||||
}
|
||||
|
||||
func TestVAPIDKeyPersistence(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
logger := slog.Default()
|
||||
store := newMockStore()
|
||||
|
||||
// First service creation should generate keys
|
||||
svc1, err := NewService(store, tmpDir, logger)
|
||||
if err != nil {
|
||||
t.Fatalf("NewService: %v", err)
|
||||
}
|
||||
key1 := svc1.GetVAPIDPublicKey()
|
||||
if key1 == "" {
|
||||
t.Fatal("VAPID public key is empty")
|
||||
}
|
||||
|
||||
// Verify file was written
|
||||
keysPath := filepath.Join(tmpDir, "vapid_keys.json")
|
||||
data, err := os.ReadFile(keysPath)
|
||||
if err != nil {
|
||||
t.Fatalf("read VAPID keys file: %v", err)
|
||||
}
|
||||
|
||||
var savedKeys VAPIDKeys
|
||||
if err := json.Unmarshal(data, &savedKeys); err != nil {
|
||||
t.Fatalf("unmarshal saved keys: %v", err)
|
||||
}
|
||||
if savedKeys.PublicKey != key1 {
|
||||
t.Errorf("saved public key = %q, service key = %q", savedKeys.PublicKey, key1)
|
||||
}
|
||||
|
||||
// Second service creation should load the same keys
|
||||
svc2, err := NewService(store, tmpDir, logger)
|
||||
if err != nil {
|
||||
t.Fatalf("NewService (second): %v", err)
|
||||
}
|
||||
key2 := svc2.GetVAPIDPublicKey()
|
||||
if key2 != key1 {
|
||||
t.Errorf("second service has different key: %q vs %q", key2, key1)
|
||||
}
|
||||
}
|
||||
|
||||
func TestVAPIDKeyFilePermissions(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
logger := slog.Default()
|
||||
store := newMockStore()
|
||||
|
||||
_, err := NewService(store, tmpDir, logger)
|
||||
if err != nil {
|
||||
t.Fatalf("NewService: %v", err)
|
||||
}
|
||||
|
||||
keysPath := filepath.Join(tmpDir, "vapid_keys.json")
|
||||
info, err := os.Stat(keysPath)
|
||||
if err != nil {
|
||||
t.Fatalf("stat VAPID keys file: %v", err)
|
||||
}
|
||||
|
||||
perm := info.Mode().Perm()
|
||||
if perm&0o077 != 0 {
|
||||
t.Errorf("VAPID keys file has overly permissive mode: %o (want 0600)", perm)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSendToUser_NoSubscriptions(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
logger := slog.Default()
|
||||
store := newMockStore()
|
||||
|
||||
svc, err := NewService(store, tmpDir, logger)
|
||||
if err != nil {
|
||||
t.Fatalf("NewService: %v", err)
|
||||
}
|
||||
|
||||
// Sending to a user with no subscriptions should not error
|
||||
err = svc.SendToUser(context.Background(), 999, Notification{
|
||||
Title: "Test",
|
||||
Body: "Test body",
|
||||
})
|
||||
if err != nil {
|
||||
t.Errorf("SendToUser with no subs: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// generateTestSubscriptionKeys creates a valid ECDSA P-256 key pair
|
||||
// encoded as base64url for use as Web Push subscription keys in tests.
|
||||
func generateTestSubscriptionKeys(t *testing.T) (p256dh, auth string) {
|
||||
t.Helper()
|
||||
privKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
||||
if err != nil {
|
||||
t.Fatalf("generate test key: %v", err)
|
||||
}
|
||||
pubBytes := elliptic.Marshal(elliptic.P256(), privKey.PublicKey.X, privKey.PublicKey.Y)
|
||||
p256dh = base64.RawURLEncoding.EncodeToString(pubBytes)
|
||||
|
||||
authBytes := make([]byte, 16)
|
||||
rand.Read(authBytes)
|
||||
auth = base64.RawURLEncoding.EncodeToString(authBytes)
|
||||
return
|
||||
}
|
||||
|
||||
func TestSendToUser_GoneSubscription(t *testing.T) {
|
||||
// Set up a fake push endpoint that returns 410 Gone
|
||||
var requestCount atomic.Int32
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
requestCount.Add(1)
|
||||
w.WriteHeader(http.StatusGone)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
tmpDir := t.TempDir()
|
||||
logger := slog.Default()
|
||||
store := newMockStore()
|
||||
|
||||
svc, err := NewService(store, tmpDir, logger)
|
||||
if err != nil {
|
||||
t.Fatalf("NewService: %v", err)
|
||||
}
|
||||
|
||||
p256dh, authKey := generateTestSubscriptionKeys(t)
|
||||
store.Subscribe(context.Background(), 1, server.URL+"/push", p256dh, authKey, "TestAgent")
|
||||
|
||||
// Send notification — should succeed but remove the stale subscription
|
||||
_ = svc.SendToUser(context.Background(), 1, Notification{
|
||||
Title: "Test",
|
||||
Body: "Test body",
|
||||
})
|
||||
|
||||
if requestCount.Load() == 0 {
|
||||
t.Fatal("expected at least one request to push endpoint")
|
||||
}
|
||||
|
||||
// The stale subscription should have been deleted
|
||||
if len(store.deleted) != 1 {
|
||||
t.Errorf("expected 1 deleted subscription, got %d", len(store.deleted))
|
||||
}
|
||||
if len(store.deleted) > 0 && store.deleted[0] != server.URL+"/push" {
|
||||
t.Errorf("deleted wrong endpoint: %q", store.deleted[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestSendToUser_NotFoundSubscription(t *testing.T) {
|
||||
// Set up a fake push endpoint that returns 404
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
tmpDir := t.TempDir()
|
||||
logger := slog.Default()
|
||||
store := newMockStore()
|
||||
|
||||
svc, err := NewService(store, tmpDir, logger)
|
||||
if err != nil {
|
||||
t.Fatalf("NewService: %v", err)
|
||||
}
|
||||
|
||||
p256dh, authKey := generateTestSubscriptionKeys(t)
|
||||
store.Subscribe(context.Background(), 1, server.URL+"/push", p256dh, authKey, "TestAgent")
|
||||
|
||||
_ = svc.SendToUser(context.Background(), 1, Notification{
|
||||
Title: "Test",
|
||||
Body: "Test body",
|
||||
})
|
||||
|
||||
// 404 should also trigger deletion
|
||||
if len(store.deleted) != 1 {
|
||||
t.Errorf("expected 1 deleted subscription, got %d", len(store.deleted))
|
||||
}
|
||||
}
|
||||
|
||||
func TestServiceSubscribeUnsubscribe(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
logger := slog.Default()
|
||||
store := newMockStore()
|
||||
|
||||
svc, err := NewService(store, tmpDir, logger)
|
||||
if err != nil {
|
||||
t.Fatalf("NewService: %v", err)
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Subscribe
|
||||
err = svc.Subscribe(ctx, 1, "https://push.example.com/test", "p256dh", "auth", "Agent")
|
||||
if err != nil {
|
||||
t.Fatalf("Subscribe: %v", err)
|
||||
}
|
||||
|
||||
subs := store.subscriptions[1]
|
||||
if len(subs) != 1 {
|
||||
t.Fatalf("got %d subscriptions, want 1", len(subs))
|
||||
}
|
||||
|
||||
// Unsubscribe
|
||||
err = svc.Unsubscribe(ctx, 1, "https://push.example.com/test")
|
||||
if err != nil {
|
||||
t.Fatalf("Unsubscribe: %v", err)
|
||||
}
|
||||
|
||||
subs = store.subscriptions[1]
|
||||
if len(subs) != 0 {
|
||||
t.Errorf("got %d subscriptions after unsubscribe, want 0", len(subs))
|
||||
}
|
||||
}
|
||||
|
||||
func TestTruncateEndpoint(t *testing.T) {
|
||||
short := "https://short.url"
|
||||
long := "https://push.services.mozilla.com/wpush/v2/very-long-subscription-id-that-goes-on-and-on"
|
||||
|
||||
got := truncateEndpoint(short)
|
||||
if got != short {
|
||||
t.Errorf("truncateEndpoint(short) = %q, want %q", got, short)
|
||||
}
|
||||
|
||||
got = truncateEndpoint(long)
|
||||
if len(got) != 40 {
|
||||
t.Errorf("truncateEndpoint(long) length = %d, want 40", len(got))
|
||||
}
|
||||
if got[len(got)-3:] != "..." {
|
||||
t.Errorf("truncateEndpoint(long) should end with '...', got %q", got)
|
||||
}
|
||||
|
||||
got = truncateEndpoint("")
|
||||
if got != "" {
|
||||
t.Errorf("truncateEndpoint('') = %q, want ''", got)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,104 @@
|
||||
// Package push provides Web Push notification support for SynapBus.
|
||||
package push
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
)
|
||||
|
||||
// Store defines the persistence interface for push subscriptions.
|
||||
type Store interface {
|
||||
// Subscribe registers a push subscription for a user.
|
||||
// If the endpoint already exists, it updates the keys and user agent.
|
||||
Subscribe(ctx context.Context, userID int64, endpoint, keyP256dh, keyAuth, userAgent string) error
|
||||
|
||||
// Unsubscribe removes a push subscription by endpoint, scoped to user.
|
||||
Unsubscribe(ctx context.Context, userID int64, endpoint string) error
|
||||
|
||||
// GetSubscriptions returns all push subscriptions for a user.
|
||||
GetSubscriptions(ctx context.Context, userID int64) ([]Subscription, error)
|
||||
|
||||
// DeleteSubscription removes a push subscription by endpoint.
|
||||
// This is used to clean up stale subscriptions (e.g., 410 Gone responses).
|
||||
DeleteSubscription(ctx context.Context, endpoint string) error
|
||||
}
|
||||
|
||||
// Subscription represents a Web Push subscription stored in the database.
|
||||
type Subscription struct {
|
||||
ID int64
|
||||
UserID int64
|
||||
Endpoint string
|
||||
KeyP256dh string
|
||||
KeyAuth string
|
||||
UserAgent string
|
||||
CreatedAt string
|
||||
}
|
||||
|
||||
// SQLiteStore implements Store using SQLite.
|
||||
type SQLiteStore struct {
|
||||
db *sql.DB
|
||||
}
|
||||
|
||||
// NewSQLiteStore creates a new SQLite-backed push subscription store.
|
||||
func NewSQLiteStore(db *sql.DB) *SQLiteStore {
|
||||
return &SQLiteStore{db: db}
|
||||
}
|
||||
|
||||
// Subscribe registers or updates a push subscription for a user.
|
||||
func (s *SQLiteStore) Subscribe(ctx context.Context, userID int64, endpoint, keyP256dh, keyAuth, userAgent string) error {
|
||||
_, err := s.db.ExecContext(ctx,
|
||||
`INSERT INTO push_subscriptions (user_id, endpoint, key_p256dh, key_auth, user_agent)
|
||||
VALUES (?, ?, ?, ?, ?)
|
||||
ON CONFLICT(endpoint) DO UPDATE SET
|
||||
key_p256dh = excluded.key_p256dh,
|
||||
key_auth = excluded.key_auth,
|
||||
user_agent = excluded.user_agent,
|
||||
user_id = excluded.user_id`,
|
||||
userID, endpoint, keyP256dh, keyAuth, userAgent,
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
// Unsubscribe removes a push subscription by endpoint, scoped to the owning user.
|
||||
func (s *SQLiteStore) Unsubscribe(ctx context.Context, userID int64, endpoint string) error {
|
||||
_, err := s.db.ExecContext(ctx,
|
||||
`DELETE FROM push_subscriptions WHERE user_id = ? AND endpoint = ?`,
|
||||
userID, endpoint,
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
// GetSubscriptions returns all push subscriptions for a user.
|
||||
func (s *SQLiteStore) GetSubscriptions(ctx context.Context, userID int64) ([]Subscription, error) {
|
||||
rows, err := s.db.QueryContext(ctx,
|
||||
`SELECT id, user_id, endpoint, key_p256dh, key_auth, user_agent, created_at
|
||||
FROM push_subscriptions
|
||||
WHERE user_id = ?
|
||||
ORDER BY created_at DESC`,
|
||||
userID,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var subs []Subscription
|
||||
for rows.Next() {
|
||||
var sub Subscription
|
||||
if err := rows.Scan(&sub.ID, &sub.UserID, &sub.Endpoint, &sub.KeyP256dh, &sub.KeyAuth, &sub.UserAgent, &sub.CreatedAt); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
subs = append(subs, sub)
|
||||
}
|
||||
return subs, rows.Err()
|
||||
}
|
||||
|
||||
// DeleteSubscription removes a push subscription by endpoint (any user).
|
||||
// Used for stale subscription cleanup (e.g., 410 Gone responses).
|
||||
func (s *SQLiteStore) DeleteSubscription(ctx context.Context, endpoint string) error {
|
||||
_, err := s.db.ExecContext(ctx,
|
||||
`DELETE FROM push_subscriptions WHERE endpoint = ?`,
|
||||
endpoint,
|
||||
)
|
||||
return err
|
||||
}
|
||||
@@ -0,0 +1,228 @@
|
||||
package push
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"testing"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
|
||||
"github.com/synapbus/synapbus/internal/storage"
|
||||
)
|
||||
|
||||
func newTestDB(t *testing.T) *sql.DB {
|
||||
t.Helper()
|
||||
dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", t.Name())
|
||||
db, err := sql.Open("sqlite", dsn)
|
||||
if err != nil {
|
||||
t.Fatalf("open database: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { db.Close() })
|
||||
|
||||
if _, err := db.Exec("PRAGMA foreign_keys=ON"); err != nil {
|
||||
t.Fatalf("enable foreign keys: %v", err)
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
if err := storage.RunMigrations(ctx, db); err != nil {
|
||||
t.Fatalf("run migrations: %v", err)
|
||||
}
|
||||
|
||||
// Seed a test user
|
||||
db.Exec(`INSERT OR IGNORE INTO users (id, username, password_hash, display_name) VALUES (1, 'testuser', 'hash', 'Test User')`)
|
||||
db.Exec(`INSERT OR IGNORE INTO users (id, username, password_hash, display_name) VALUES (2, 'otheruser', 'hash', 'Other User')`)
|
||||
|
||||
return db
|
||||
}
|
||||
|
||||
func TestSQLiteStore_Subscribe(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
err := store.Subscribe(ctx, 1, "https://push.example.com/sub1", "p256dh-key-1", "auth-key-1", "TestAgent/1.0")
|
||||
if err != nil {
|
||||
t.Fatalf("Subscribe: %v", err)
|
||||
}
|
||||
|
||||
subs, err := store.GetSubscriptions(ctx, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("GetSubscriptions: %v", err)
|
||||
}
|
||||
if len(subs) != 1 {
|
||||
t.Fatalf("got %d subscriptions, want 1", len(subs))
|
||||
}
|
||||
|
||||
sub := subs[0]
|
||||
if sub.Endpoint != "https://push.example.com/sub1" {
|
||||
t.Errorf("endpoint = %q, want %q", sub.Endpoint, "https://push.example.com/sub1")
|
||||
}
|
||||
if sub.KeyP256dh != "p256dh-key-1" {
|
||||
t.Errorf("key_p256dh = %q, want %q", sub.KeyP256dh, "p256dh-key-1")
|
||||
}
|
||||
if sub.KeyAuth != "auth-key-1" {
|
||||
t.Errorf("key_auth = %q, want %q", sub.KeyAuth, "auth-key-1")
|
||||
}
|
||||
if sub.UserAgent != "TestAgent/1.0" {
|
||||
t.Errorf("user_agent = %q, want %q", sub.UserAgent, "TestAgent/1.0")
|
||||
}
|
||||
if sub.UserID != 1 {
|
||||
t.Errorf("user_id = %d, want 1", sub.UserID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteStore_SubscribeUpsert(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
// First subscription
|
||||
err := store.Subscribe(ctx, 1, "https://push.example.com/sub1", "old-p256dh", "old-auth", "OldAgent")
|
||||
if err != nil {
|
||||
t.Fatalf("Subscribe: %v", err)
|
||||
}
|
||||
|
||||
// Update same endpoint with new keys
|
||||
err = store.Subscribe(ctx, 1, "https://push.example.com/sub1", "new-p256dh", "new-auth", "NewAgent")
|
||||
if err != nil {
|
||||
t.Fatalf("Subscribe upsert: %v", err)
|
||||
}
|
||||
|
||||
subs, err := store.GetSubscriptions(ctx, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("GetSubscriptions: %v", err)
|
||||
}
|
||||
if len(subs) != 1 {
|
||||
t.Fatalf("got %d subscriptions, want 1 (should upsert)", len(subs))
|
||||
}
|
||||
if subs[0].KeyP256dh != "new-p256dh" {
|
||||
t.Errorf("key_p256dh = %q, want %q", subs[0].KeyP256dh, "new-p256dh")
|
||||
}
|
||||
if subs[0].KeyAuth != "new-auth" {
|
||||
t.Errorf("key_auth = %q, want %q", subs[0].KeyAuth, "new-auth")
|
||||
}
|
||||
if subs[0].UserAgent != "NewAgent" {
|
||||
t.Errorf("user_agent = %q, want %q", subs[0].UserAgent, "NewAgent")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteStore_MultipleSubscriptions(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
// Register multiple subscriptions for same user
|
||||
for i := 0; i < 3; i++ {
|
||||
endpoint := fmt.Sprintf("https://push.example.com/sub%d", i)
|
||||
err := store.Subscribe(ctx, 1, endpoint, fmt.Sprintf("p256dh-%d", i), fmt.Sprintf("auth-%d", i), "Agent")
|
||||
if err != nil {
|
||||
t.Fatalf("Subscribe(%d): %v", i, err)
|
||||
}
|
||||
}
|
||||
|
||||
subs, err := store.GetSubscriptions(ctx, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("GetSubscriptions: %v", err)
|
||||
}
|
||||
if len(subs) != 3 {
|
||||
t.Errorf("got %d subscriptions, want 3", len(subs))
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteStore_Unsubscribe(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
store.Subscribe(ctx, 1, "https://push.example.com/sub1", "p256dh", "auth", "Agent")
|
||||
store.Subscribe(ctx, 1, "https://push.example.com/sub2", "p256dh2", "auth2", "Agent")
|
||||
|
||||
err := store.Unsubscribe(ctx, 1, "https://push.example.com/sub1")
|
||||
if err != nil {
|
||||
t.Fatalf("Unsubscribe: %v", err)
|
||||
}
|
||||
|
||||
subs, err := store.GetSubscriptions(ctx, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("GetSubscriptions: %v", err)
|
||||
}
|
||||
if len(subs) != 1 {
|
||||
t.Fatalf("got %d subscriptions, want 1", len(subs))
|
||||
}
|
||||
if subs[0].Endpoint != "https://push.example.com/sub2" {
|
||||
t.Errorf("remaining endpoint = %q, want sub2", subs[0].Endpoint)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteStore_DeleteSubscription(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
store.Subscribe(ctx, 1, "https://push.example.com/sub1", "p256dh", "auth", "Agent")
|
||||
|
||||
err := store.DeleteSubscription(ctx, "https://push.example.com/sub1")
|
||||
if err != nil {
|
||||
t.Fatalf("DeleteSubscription: %v", err)
|
||||
}
|
||||
|
||||
subs, err := store.GetSubscriptions(ctx, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("GetSubscriptions: %v", err)
|
||||
}
|
||||
if len(subs) != 0 {
|
||||
t.Errorf("got %d subscriptions, want 0", len(subs))
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteStore_UserIsolation(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
store.Subscribe(ctx, 1, "https://push.example.com/user1-sub", "p256dh1", "auth1", "Agent")
|
||||
store.Subscribe(ctx, 2, "https://push.example.com/user2-sub", "p256dh2", "auth2", "Agent")
|
||||
|
||||
subs1, err := store.GetSubscriptions(ctx, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("GetSubscriptions(1): %v", err)
|
||||
}
|
||||
if len(subs1) != 1 {
|
||||
t.Errorf("user 1: got %d subscriptions, want 1", len(subs1))
|
||||
}
|
||||
|
||||
subs2, err := store.GetSubscriptions(ctx, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("GetSubscriptions(2): %v", err)
|
||||
}
|
||||
if len(subs2) != 1 {
|
||||
t.Errorf("user 2: got %d subscriptions, want 1", len(subs2))
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteStore_GetSubscriptionsEmpty(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
subs, err := store.GetSubscriptions(ctx, 999)
|
||||
if err != nil {
|
||||
t.Fatalf("GetSubscriptions: %v", err)
|
||||
}
|
||||
if subs != nil {
|
||||
t.Errorf("got %v, want nil for empty result", subs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteStore_UnsubscribeNonexistent(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
// Should not error when deleting non-existent endpoint
|
||||
err := store.Unsubscribe(ctx, 1, "https://push.example.com/nonexistent")
|
||||
if err != nil {
|
||||
t.Fatalf("Unsubscribe non-existent: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,113 @@
|
||||
// Package reactions provides message reaction types, storage, and workflow state logic.
|
||||
package reactions
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Valid reaction types.
|
||||
const (
|
||||
ReactionApprove = "approve"
|
||||
ReactionReject = "reject"
|
||||
ReactionInProgress = "in_progress"
|
||||
ReactionDone = "done"
|
||||
ReactionPublished = "published"
|
||||
)
|
||||
|
||||
// Workflow states (derived from reactions).
|
||||
const (
|
||||
StateProposed = "proposed"
|
||||
StateApproved = "approved"
|
||||
StateInProgress = "in_progress"
|
||||
StateRejected = "rejected"
|
||||
StateDone = "done"
|
||||
StatePublished = "published"
|
||||
)
|
||||
|
||||
// reactionPriority maps reaction types to their priority for state derivation.
|
||||
// Higher number = higher priority = wins for badge display.
|
||||
var reactionPriority = map[string]int{
|
||||
ReactionApprove: 2,
|
||||
ReactionInProgress: 3,
|
||||
ReactionReject: 4,
|
||||
ReactionDone: 5,
|
||||
ReactionPublished: 6,
|
||||
}
|
||||
|
||||
// reactionToState maps reaction types to workflow states.
|
||||
var reactionToState = map[string]string{
|
||||
ReactionApprove: StateApproved,
|
||||
ReactionReject: StateRejected,
|
||||
ReactionInProgress: StateInProgress,
|
||||
ReactionDone: StateDone,
|
||||
ReactionPublished: StatePublished,
|
||||
}
|
||||
|
||||
// TerminalStates are states that should not trigger stalemate checks.
|
||||
var TerminalStates = map[string]bool{
|
||||
StateRejected: true,
|
||||
StateDone: true,
|
||||
StatePublished: true,
|
||||
}
|
||||
|
||||
// MaxReactionsPerMessage is the safety limit.
|
||||
const MaxReactionsPerMessage = 100
|
||||
|
||||
// Sentinel errors.
|
||||
var (
|
||||
ErrInvalidReaction = errors.New("invalid reaction type: must be one of approve, reject, in_progress, done, published")
|
||||
ErrReactionLimit = errors.New("maximum reactions per message (100) reached")
|
||||
ErrNotMember = errors.New("only channel members can react to messages")
|
||||
)
|
||||
|
||||
// Reaction represents a single reaction on a message.
|
||||
type Reaction struct {
|
||||
ID int64 `json:"id"`
|
||||
MessageID int64 `json:"message_id"`
|
||||
AgentName string `json:"agent_name"`
|
||||
Reaction string `json:"reaction"`
|
||||
Metadata json.RawMessage `json:"metadata"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
|
||||
// ValidReactions is the set of allowed reaction types.
|
||||
var ValidReactions = map[string]bool{
|
||||
ReactionApprove: true,
|
||||
ReactionReject: true,
|
||||
ReactionInProgress: true,
|
||||
ReactionDone: true,
|
||||
ReactionPublished: true,
|
||||
}
|
||||
|
||||
// IsValidReaction returns true if the reaction type is valid.
|
||||
func IsValidReaction(r string) bool {
|
||||
return ValidReactions[r]
|
||||
}
|
||||
|
||||
// ComputeWorkflowState derives the workflow state from a list of reactions.
|
||||
// Returns "proposed" if no reactions exist.
|
||||
func ComputeWorkflowState(reactions []*Reaction) string {
|
||||
if len(reactions) == 0 {
|
||||
return StateProposed
|
||||
}
|
||||
|
||||
highestPriority := 0
|
||||
highestState := StateProposed
|
||||
|
||||
for _, r := range reactions {
|
||||
p, ok := reactionPriority[r.Reaction]
|
||||
if ok && p > highestPriority {
|
||||
highestPriority = p
|
||||
highestState = reactionToState[r.Reaction]
|
||||
}
|
||||
}
|
||||
|
||||
return highestState
|
||||
}
|
||||
|
||||
// IsTerminalState returns true if the state should not trigger stalemate checks.
|
||||
func IsTerminalState(state string) bool {
|
||||
return TerminalStates[state]
|
||||
}
|
||||
@@ -0,0 +1,150 @@
|
||||
package reactions
|
||||
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestIsValidReaction(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
reaction string
|
||||
want bool
|
||||
}{
|
||||
{"approve is valid", ReactionApprove, true},
|
||||
{"reject is valid", ReactionReject, true},
|
||||
{"in_progress is valid", ReactionInProgress, true},
|
||||
{"done is valid", ReactionDone, true},
|
||||
{"published is valid", ReactionPublished, true},
|
||||
{"empty string is invalid", "", false},
|
||||
{"thumbs_up is invalid", "thumbs_up", false},
|
||||
{"like is invalid", "like", false},
|
||||
{"APPROVE uppercase is invalid", "APPROVE", false},
|
||||
{"random text is invalid", "foobar", false},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got := IsValidReaction(tt.reaction)
|
||||
if got != tt.want {
|
||||
t.Errorf("IsValidReaction(%q) = %v, want %v", tt.reaction, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestComputeWorkflowState(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
reactions []*Reaction
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "empty reactions returns proposed",
|
||||
reactions: []*Reaction{},
|
||||
want: StateProposed,
|
||||
},
|
||||
{
|
||||
name: "nil reactions returns proposed",
|
||||
reactions: nil,
|
||||
want: StateProposed,
|
||||
},
|
||||
{
|
||||
name: "single approve returns approved",
|
||||
reactions: []*Reaction{
|
||||
{Reaction: ReactionApprove, AgentName: "agent-a"},
|
||||
},
|
||||
want: StateApproved,
|
||||
},
|
||||
{
|
||||
name: "approve + in_progress returns in_progress (higher priority wins)",
|
||||
reactions: []*Reaction{
|
||||
{Reaction: ReactionApprove, AgentName: "agent-a"},
|
||||
{Reaction: ReactionInProgress, AgentName: "agent-b"},
|
||||
},
|
||||
want: StateInProgress,
|
||||
},
|
||||
{
|
||||
name: "single reject returns rejected",
|
||||
reactions: []*Reaction{
|
||||
{Reaction: ReactionReject, AgentName: "agent-a"},
|
||||
},
|
||||
want: StateRejected,
|
||||
},
|
||||
{
|
||||
name: "single done returns done",
|
||||
reactions: []*Reaction{
|
||||
{Reaction: ReactionDone, AgentName: "agent-a"},
|
||||
},
|
||||
want: StateDone,
|
||||
},
|
||||
{
|
||||
name: "single published returns published",
|
||||
reactions: []*Reaction{
|
||||
{Reaction: ReactionPublished, AgentName: "agent-a"},
|
||||
},
|
||||
want: StatePublished,
|
||||
},
|
||||
{
|
||||
name: "all five types - published wins",
|
||||
reactions: []*Reaction{
|
||||
{Reaction: ReactionApprove, AgentName: "agent-a"},
|
||||
{Reaction: ReactionInProgress, AgentName: "agent-b"},
|
||||
{Reaction: ReactionReject, AgentName: "agent-c"},
|
||||
{Reaction: ReactionDone, AgentName: "agent-d"},
|
||||
{Reaction: ReactionPublished, AgentName: "agent-e"},
|
||||
},
|
||||
want: StatePublished,
|
||||
},
|
||||
{
|
||||
name: "published wins over everything",
|
||||
reactions: []*Reaction{
|
||||
{Reaction: ReactionDone, AgentName: "agent-a"},
|
||||
{Reaction: ReactionReject, AgentName: "agent-b"},
|
||||
{Reaction: ReactionPublished, AgentName: "agent-c"},
|
||||
},
|
||||
want: StatePublished,
|
||||
},
|
||||
{
|
||||
name: "reject beats in_progress",
|
||||
reactions: []*Reaction{
|
||||
{Reaction: ReactionInProgress, AgentName: "agent-a"},
|
||||
{Reaction: ReactionReject, AgentName: "agent-b"},
|
||||
},
|
||||
want: StateRejected,
|
||||
},
|
||||
{
|
||||
name: "done beats reject",
|
||||
reactions: []*Reaction{
|
||||
{Reaction: ReactionReject, AgentName: "agent-a"},
|
||||
{Reaction: ReactionDone, AgentName: "agent-b"},
|
||||
},
|
||||
want: StateDone,
|
||||
},
|
||||
{
|
||||
name: "in_progress beats approve",
|
||||
reactions: []*Reaction{
|
||||
{Reaction: ReactionInProgress, AgentName: "agent-a"},
|
||||
{Reaction: ReactionApprove, AgentName: "agent-b"},
|
||||
},
|
||||
want: StateInProgress,
|
||||
},
|
||||
{
|
||||
name: "multiple approves still returns approved",
|
||||
reactions: []*Reaction{
|
||||
{Reaction: ReactionApprove, AgentName: "agent-a"},
|
||||
{Reaction: ReactionApprove, AgentName: "agent-b"},
|
||||
{Reaction: ReactionApprove, AgentName: "agent-c"},
|
||||
},
|
||||
want: StateApproved,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got := ComputeWorkflowState(tt.reactions)
|
||||
if got != tt.want {
|
||||
t.Errorf("ComputeWorkflowState() = %q, want %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,127 @@
|
||||
package reactions
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
)
|
||||
|
||||
// Service provides business logic for message reactions.
|
||||
type Service struct {
|
||||
store Store
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
// NewService creates a new reaction service.
|
||||
func NewService(store Store, logger *slog.Logger) *Service {
|
||||
return &Service{
|
||||
store: store,
|
||||
logger: logger.With("component", "reactions"),
|
||||
}
|
||||
}
|
||||
|
||||
// ToggleResult describes what happened after a toggle operation.
|
||||
type ToggleResult struct {
|
||||
Action string `json:"action"` // "added" or "removed"
|
||||
Reaction *Reaction `json:"reaction,omitempty"`
|
||||
}
|
||||
|
||||
// Toggle adds a reaction if it doesn't exist, or removes it if it does.
|
||||
// Returns the action taken and the reaction (if added).
|
||||
func (s *Service) Toggle(ctx context.Context, messageID int64, agentName, reactionType string, metadata json.RawMessage) (*ToggleResult, error) {
|
||||
if !IsValidReaction(reactionType) {
|
||||
return nil, ErrInvalidReaction
|
||||
}
|
||||
|
||||
// Check if reaction already exists
|
||||
exists, err := s.store.Exists(ctx, messageID, agentName, reactionType)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("check existing reaction: %w", err)
|
||||
}
|
||||
|
||||
if exists {
|
||||
// Toggle off — remove it
|
||||
if err := s.store.Delete(ctx, messageID, agentName, reactionType); err != nil {
|
||||
return nil, fmt.Errorf("remove reaction: %w", err)
|
||||
}
|
||||
s.logger.Info("reaction removed",
|
||||
"message_id", messageID,
|
||||
"agent", agentName,
|
||||
"reaction", reactionType,
|
||||
)
|
||||
return &ToggleResult{Action: "removed"}, nil
|
||||
}
|
||||
|
||||
// Check reaction count limit
|
||||
count, err := s.store.CountByMessage(ctx, messageID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("count reactions: %w", err)
|
||||
}
|
||||
if count >= MaxReactionsPerMessage {
|
||||
return nil, ErrReactionLimit
|
||||
}
|
||||
|
||||
// Toggle on — add it
|
||||
if metadata == nil {
|
||||
metadata = json.RawMessage("{}")
|
||||
}
|
||||
|
||||
r := &Reaction{
|
||||
MessageID: messageID,
|
||||
AgentName: agentName,
|
||||
Reaction: reactionType,
|
||||
Metadata: metadata,
|
||||
}
|
||||
|
||||
if err := s.store.Insert(ctx, r); err != nil {
|
||||
return nil, fmt.Errorf("add reaction: %w", err)
|
||||
}
|
||||
|
||||
s.logger.Info("reaction added",
|
||||
"message_id", messageID,
|
||||
"agent", agentName,
|
||||
"reaction", reactionType,
|
||||
)
|
||||
|
||||
return &ToggleResult{Action: "added", Reaction: r}, nil
|
||||
}
|
||||
|
||||
// Remove explicitly removes a reaction.
|
||||
func (s *Service) Remove(ctx context.Context, messageID int64, agentName, reactionType string) error {
|
||||
if !IsValidReaction(reactionType) {
|
||||
return ErrInvalidReaction
|
||||
}
|
||||
|
||||
if err := s.store.Delete(ctx, messageID, agentName, reactionType); err != nil {
|
||||
return fmt.Errorf("remove reaction: %w", err)
|
||||
}
|
||||
|
||||
s.logger.Info("reaction removed",
|
||||
"message_id", messageID,
|
||||
"agent", agentName,
|
||||
"reaction", reactionType,
|
||||
)
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetReactions returns all reactions for a message and the computed workflow state.
|
||||
func (s *Service) GetReactions(ctx context.Context, messageID int64) ([]*Reaction, string, error) {
|
||||
reactions, err := s.store.GetByMessageID(ctx, messageID)
|
||||
if err != nil {
|
||||
return nil, "", fmt.Errorf("get reactions: %w", err)
|
||||
}
|
||||
|
||||
state := ComputeWorkflowState(reactions)
|
||||
return reactions, state, nil
|
||||
}
|
||||
|
||||
// GetReactionsByMessageIDs returns reactions grouped by message ID.
|
||||
func (s *Service) GetReactionsByMessageIDs(ctx context.Context, messageIDs []int64) (map[int64][]*Reaction, error) {
|
||||
return s.store.GetByMessageIDs(ctx, messageIDs)
|
||||
}
|
||||
|
||||
// ListByState returns message IDs in a channel that have the given workflow state.
|
||||
func (s *Service) ListByState(ctx context.Context, channelID int64, state string) ([]int64, error) {
|
||||
return s.store.GetMessageIDsByState(ctx, channelID, state)
|
||||
}
|
||||
@@ -0,0 +1,209 @@
|
||||
package reactions
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// Store defines the storage interface for reactions.
|
||||
type Store interface {
|
||||
Insert(ctx context.Context, r *Reaction) error
|
||||
Delete(ctx context.Context, messageID int64, agentName, reaction string) error
|
||||
GetByMessageID(ctx context.Context, messageID int64) ([]*Reaction, error)
|
||||
GetByMessageIDs(ctx context.Context, messageIDs []int64) (map[int64][]*Reaction, error)
|
||||
Exists(ctx context.Context, messageID int64, agentName, reaction string) (bool, error)
|
||||
CountByMessage(ctx context.Context, messageID int64) (int, error)
|
||||
// GetMessageIDsByState returns message IDs in a channel that have the given workflow state.
|
||||
GetMessageIDsByState(ctx context.Context, channelID int64, state string) ([]int64, error)
|
||||
}
|
||||
|
||||
// SQLiteStore implements Store using SQLite.
|
||||
type SQLiteStore struct {
|
||||
db *sql.DB
|
||||
}
|
||||
|
||||
// NewSQLiteStore creates a new SQLite-backed reaction store.
|
||||
func NewSQLiteStore(db *sql.DB) *SQLiteStore {
|
||||
return &SQLiteStore{db: db}
|
||||
}
|
||||
|
||||
func (s *SQLiteStore) Insert(ctx context.Context, r *Reaction) error {
|
||||
metadata := r.Metadata
|
||||
if metadata == nil {
|
||||
metadata = json.RawMessage("{}")
|
||||
}
|
||||
|
||||
result, err := s.db.ExecContext(ctx,
|
||||
`INSERT INTO message_reactions (message_id, agent_name, reaction, metadata, created_at)
|
||||
VALUES (?, ?, ?, ?, CURRENT_TIMESTAMP)`,
|
||||
r.MessageID, r.AgentName, r.Reaction, string(metadata),
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("insert reaction: %w", err)
|
||||
}
|
||||
id, err := result.LastInsertId()
|
||||
if err != nil {
|
||||
return fmt.Errorf("get reaction id: %w", err)
|
||||
}
|
||||
r.ID = id
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *SQLiteStore) Delete(ctx context.Context, messageID int64, agentName, reaction string) error {
|
||||
result, err := s.db.ExecContext(ctx,
|
||||
`DELETE FROM message_reactions WHERE message_id = ? AND agent_name = ? AND reaction = ?`,
|
||||
messageID, agentName, reaction,
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("delete reaction: %w", err)
|
||||
}
|
||||
n, _ := result.RowsAffected()
|
||||
if n == 0 {
|
||||
return fmt.Errorf("reaction not found")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *SQLiteStore) GetByMessageID(ctx context.Context, messageID int64) ([]*Reaction, error) {
|
||||
rows, err := s.db.QueryContext(ctx,
|
||||
`SELECT id, message_id, agent_name, reaction, metadata, created_at
|
||||
FROM message_reactions WHERE message_id = ?
|
||||
ORDER BY created_at ASC`, messageID,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get reactions: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
return scanReactions(rows)
|
||||
}
|
||||
|
||||
func (s *SQLiteStore) GetByMessageIDs(ctx context.Context, messageIDs []int64) (map[int64][]*Reaction, error) {
|
||||
if len(messageIDs) == 0 {
|
||||
return map[int64][]*Reaction{}, nil
|
||||
}
|
||||
|
||||
placeholders := make([]string, len(messageIDs))
|
||||
args := make([]any, len(messageIDs))
|
||||
for i, id := range messageIDs {
|
||||
placeholders[i] = "?"
|
||||
args[i] = id
|
||||
}
|
||||
|
||||
query := fmt.Sprintf(
|
||||
`SELECT id, message_id, agent_name, reaction, metadata, created_at
|
||||
FROM message_reactions WHERE message_id IN (%s)
|
||||
ORDER BY created_at ASC`,
|
||||
strings.Join(placeholders, ","),
|
||||
)
|
||||
|
||||
rows, err := s.db.QueryContext(ctx, query, args...)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get reactions by ids: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
all, err := scanReactions(rows)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
result := make(map[int64][]*Reaction)
|
||||
for _, r := range all {
|
||||
result[r.MessageID] = append(result[r.MessageID], r)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (s *SQLiteStore) Exists(ctx context.Context, messageID int64, agentName, reaction string) (bool, error) {
|
||||
var count int
|
||||
err := s.db.QueryRowContext(ctx,
|
||||
`SELECT COUNT(*) FROM message_reactions WHERE message_id = ? AND agent_name = ? AND reaction = ?`,
|
||||
messageID, agentName, reaction,
|
||||
).Scan(&count)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("check reaction exists: %w", err)
|
||||
}
|
||||
return count > 0, nil
|
||||
}
|
||||
|
||||
func (s *SQLiteStore) CountByMessage(ctx context.Context, messageID int64) (int, error) {
|
||||
var count int
|
||||
err := s.db.QueryRowContext(ctx,
|
||||
`SELECT COUNT(*) FROM message_reactions WHERE message_id = ?`, messageID,
|
||||
).Scan(&count)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("count reactions: %w", err)
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
|
||||
func (s *SQLiteStore) GetMessageIDsByState(ctx context.Context, channelID int64, state string) ([]int64, error) {
|
||||
var query string
|
||||
var args []any
|
||||
|
||||
if state == StateProposed {
|
||||
// Messages with no reactions
|
||||
query = `SELECT m.id FROM messages m
|
||||
WHERE m.channel_id = ?
|
||||
AND NOT EXISTS (SELECT 1 FROM message_reactions r WHERE r.message_id = m.id)
|
||||
ORDER BY m.created_at DESC`
|
||||
args = []any{channelID}
|
||||
} else {
|
||||
// Find the reaction type for this state
|
||||
var reactionType string
|
||||
for rt, st := range reactionToState {
|
||||
if st == state {
|
||||
reactionType = rt
|
||||
break
|
||||
}
|
||||
}
|
||||
if reactionType == "" {
|
||||
return nil, fmt.Errorf("unknown workflow state: %s", state)
|
||||
}
|
||||
|
||||
// Messages where the highest-priority reaction maps to this state
|
||||
// We get all messages with this reaction type and filter in app layer
|
||||
query = `SELECT DISTINCT r.message_id FROM message_reactions r
|
||||
JOIN messages m ON m.id = r.message_id
|
||||
WHERE m.channel_id = ? AND r.reaction = ?
|
||||
ORDER BY m.created_at DESC`
|
||||
args = []any{channelID, reactionType}
|
||||
}
|
||||
|
||||
rows, err := s.db.QueryContext(ctx, query, args...)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get message ids by state: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var ids []int64
|
||||
for rows.Next() {
|
||||
var id int64
|
||||
if err := rows.Scan(&id); err != nil {
|
||||
return nil, fmt.Errorf("scan message id: %w", err)
|
||||
}
|
||||
ids = append(ids, id)
|
||||
}
|
||||
return ids, rows.Err()
|
||||
}
|
||||
|
||||
func scanReactions(rows *sql.Rows) ([]*Reaction, error) {
|
||||
var reactions []*Reaction
|
||||
for rows.Next() {
|
||||
var r Reaction
|
||||
var metadata string
|
||||
err := rows.Scan(&r.ID, &r.MessageID, &r.AgentName, &r.Reaction, &metadata, &r.CreatedAt)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("scan reaction: %w", err)
|
||||
}
|
||||
r.Metadata = json.RawMessage(metadata)
|
||||
reactions = append(reactions, &r)
|
||||
}
|
||||
if reactions == nil {
|
||||
reactions = []*Reaction{}
|
||||
}
|
||||
return reactions, rows.Err()
|
||||
}
|
||||
@@ -0,0 +1,332 @@
|
||||
package reactions
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"testing"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
|
||||
"github.com/synapbus/synapbus/internal/storage"
|
||||
)
|
||||
|
||||
func newTestDB(t *testing.T) *sql.DB {
|
||||
t.Helper()
|
||||
dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", t.Name())
|
||||
db, err := sql.Open("sqlite", dsn)
|
||||
if err != nil {
|
||||
t.Fatalf("open database: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { db.Close() })
|
||||
|
||||
if _, err := db.Exec("PRAGMA foreign_keys=ON"); err != nil {
|
||||
t.Fatalf("enable foreign keys: %v", err)
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
if err := storage.RunMigrations(ctx, db); err != nil {
|
||||
t.Fatalf("run migrations: %v", err)
|
||||
}
|
||||
|
||||
return db
|
||||
}
|
||||
|
||||
// seedTestMessage creates a test user, agent, conversation, and message,
|
||||
// returning the message ID.
|
||||
func seedTestMessage(t *testing.T, db *sql.DB, agentName string) int64 {
|
||||
t.Helper()
|
||||
|
||||
// Ensure user exists
|
||||
db.Exec(`INSERT OR IGNORE INTO users (id, username, password_hash, display_name) VALUES (1, 'testowner', 'hash', 'Test Owner')`)
|
||||
|
||||
// Ensure agent exists
|
||||
_, err := db.Exec(
|
||||
`INSERT OR IGNORE INTO agents (name, display_name, type, capabilities, owner_id, api_key_hash, status) VALUES (?, ?, 'ai', '{}', 1, 'testhash', 'active')`,
|
||||
agentName, agentName,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("seed agent %s: %v", agentName, err)
|
||||
}
|
||||
|
||||
// Create conversation
|
||||
result, err := db.Exec(
|
||||
`INSERT INTO conversations (subject, created_by, created_at, updated_at) VALUES ('test', ?, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`,
|
||||
agentName,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("create conversation: %v", err)
|
||||
}
|
||||
convID, _ := result.LastInsertId()
|
||||
|
||||
// Create message
|
||||
result, err = db.Exec(
|
||||
`INSERT INTO messages (conversation_id, from_agent, body, priority, status, created_at) VALUES (?, ?, 'test body', 5, 'pending', CURRENT_TIMESTAMP)`,
|
||||
convID, agentName,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("create message: %v", err)
|
||||
}
|
||||
msgID, _ := result.LastInsertId()
|
||||
return msgID
|
||||
}
|
||||
|
||||
func TestSQLiteStore_InsertAndGet(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
msgID := seedTestMessage(t, db, "agent-a")
|
||||
|
||||
r := &Reaction{
|
||||
MessageID: msgID,
|
||||
AgentName: "agent-a",
|
||||
Reaction: ReactionApprove,
|
||||
Metadata: json.RawMessage(`{"comment":"looks good"}`),
|
||||
}
|
||||
|
||||
if err := store.Insert(ctx, r); err != nil {
|
||||
t.Fatalf("Insert: %v", err)
|
||||
}
|
||||
|
||||
if r.ID == 0 {
|
||||
t.Error("reaction ID should not be 0 after insert")
|
||||
}
|
||||
|
||||
// Verify it exists
|
||||
exists, err := store.Exists(ctx, msgID, "agent-a", ReactionApprove)
|
||||
if err != nil {
|
||||
t.Fatalf("Exists: %v", err)
|
||||
}
|
||||
if !exists {
|
||||
t.Error("expected reaction to exist after insert")
|
||||
}
|
||||
|
||||
// GetByMessageID
|
||||
reactions, err := store.GetByMessageID(ctx, msgID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetByMessageID: %v", err)
|
||||
}
|
||||
if len(reactions) != 1 {
|
||||
t.Fatalf("got %d reactions, want 1", len(reactions))
|
||||
}
|
||||
if reactions[0].AgentName != "agent-a" {
|
||||
t.Errorf("AgentName = %q, want %q", reactions[0].AgentName, "agent-a")
|
||||
}
|
||||
if reactions[0].Reaction != ReactionApprove {
|
||||
t.Errorf("Reaction = %q, want %q", reactions[0].Reaction, ReactionApprove)
|
||||
}
|
||||
if string(reactions[0].Metadata) != `{"comment":"looks good"}` {
|
||||
t.Errorf("Metadata = %s, want %s", reactions[0].Metadata, `{"comment":"looks good"}`)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteStore_UniqueConstraint(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
msgID := seedTestMessage(t, db, "agent-a")
|
||||
|
||||
r := &Reaction{
|
||||
MessageID: msgID,
|
||||
AgentName: "agent-a",
|
||||
Reaction: ReactionApprove,
|
||||
}
|
||||
|
||||
if err := store.Insert(ctx, r); err != nil {
|
||||
t.Fatalf("Insert first: %v", err)
|
||||
}
|
||||
|
||||
// Inserting the same reaction again should fail with UNIQUE constraint
|
||||
r2 := &Reaction{
|
||||
MessageID: msgID,
|
||||
AgentName: "agent-a",
|
||||
Reaction: ReactionApprove,
|
||||
}
|
||||
err := store.Insert(ctx, r2)
|
||||
if err == nil {
|
||||
t.Error("expected error on duplicate insert, got nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteStore_Delete(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
msgID := seedTestMessage(t, db, "agent-a")
|
||||
|
||||
r := &Reaction{
|
||||
MessageID: msgID,
|
||||
AgentName: "agent-a",
|
||||
Reaction: ReactionApprove,
|
||||
}
|
||||
if err := store.Insert(ctx, r); err != nil {
|
||||
t.Fatalf("Insert: %v", err)
|
||||
}
|
||||
|
||||
// Delete the reaction
|
||||
if err := store.Delete(ctx, msgID, "agent-a", ReactionApprove); err != nil {
|
||||
t.Fatalf("Delete: %v", err)
|
||||
}
|
||||
|
||||
// Verify it's gone
|
||||
exists, err := store.Exists(ctx, msgID, "agent-a", ReactionApprove)
|
||||
if err != nil {
|
||||
t.Fatalf("Exists: %v", err)
|
||||
}
|
||||
if exists {
|
||||
t.Error("expected reaction to not exist after delete")
|
||||
}
|
||||
|
||||
// Deleting a non-existent reaction should return an error
|
||||
err = store.Delete(ctx, msgID, "agent-a", ReactionApprove)
|
||||
if err == nil {
|
||||
t.Error("expected error when deleting non-existent reaction, got nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteStore_GetByMessageID(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
msgID := seedTestMessage(t, db, "agent-a")
|
||||
// Seed a second agent
|
||||
db.Exec(`INSERT OR IGNORE INTO agents (name, display_name, type, capabilities, owner_id, api_key_hash, status) VALUES ('agent-b', 'agent-b', 'ai', '{}', 1, 'testhash2', 'active')`)
|
||||
|
||||
// Insert multiple reactions from different agents
|
||||
reactions := []*Reaction{
|
||||
{MessageID: msgID, AgentName: "agent-a", Reaction: ReactionApprove},
|
||||
{MessageID: msgID, AgentName: "agent-b", Reaction: ReactionInProgress},
|
||||
{MessageID: msgID, AgentName: "agent-a", Reaction: ReactionDone},
|
||||
}
|
||||
|
||||
for _, r := range reactions {
|
||||
if err := store.Insert(ctx, r); err != nil {
|
||||
t.Fatalf("Insert: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
got, err := store.GetByMessageID(ctx, msgID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetByMessageID: %v", err)
|
||||
}
|
||||
|
||||
if len(got) != 3 {
|
||||
t.Fatalf("got %d reactions, want 3", len(got))
|
||||
}
|
||||
|
||||
// Verify results are ordered by created_at ASC
|
||||
for i, r := range got {
|
||||
if r.ID == 0 {
|
||||
t.Errorf("reaction[%d] ID should not be 0", i)
|
||||
}
|
||||
if r.MessageID != msgID {
|
||||
t.Errorf("reaction[%d] MessageID = %d, want %d", i, r.MessageID, msgID)
|
||||
}
|
||||
}
|
||||
|
||||
// Test with a message that has no reactions
|
||||
msgID2 := seedTestMessage(t, db, "agent-a")
|
||||
got2, err := store.GetByMessageID(ctx, msgID2)
|
||||
if err != nil {
|
||||
t.Fatalf("GetByMessageID (empty): %v", err)
|
||||
}
|
||||
if len(got2) != 0 {
|
||||
t.Errorf("got %d reactions for empty message, want 0", len(got2))
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteStore_CountByMessage(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
msgID := seedTestMessage(t, db, "agent-a")
|
||||
db.Exec(`INSERT OR IGNORE INTO agents (name, display_name, type, capabilities, owner_id, api_key_hash, status) VALUES ('agent-b', 'agent-b', 'ai', '{}', 1, 'testhash2', 'active')`)
|
||||
|
||||
// Count should be 0 initially
|
||||
count, err := store.CountByMessage(ctx, msgID)
|
||||
if err != nil {
|
||||
t.Fatalf("CountByMessage: %v", err)
|
||||
}
|
||||
if count != 0 {
|
||||
t.Errorf("initial count = %d, want 0", count)
|
||||
}
|
||||
|
||||
// Insert some reactions
|
||||
for _, r := range []*Reaction{
|
||||
{MessageID: msgID, AgentName: "agent-a", Reaction: ReactionApprove},
|
||||
{MessageID: msgID, AgentName: "agent-b", Reaction: ReactionDone},
|
||||
} {
|
||||
if err := store.Insert(ctx, r); err != nil {
|
||||
t.Fatalf("Insert: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
count, err = store.CountByMessage(ctx, msgID)
|
||||
if err != nil {
|
||||
t.Fatalf("CountByMessage: %v", err)
|
||||
}
|
||||
if count != 2 {
|
||||
t.Errorf("count = %d, want 2", count)
|
||||
}
|
||||
|
||||
// Delete one and verify count decreases
|
||||
if err := store.Delete(ctx, msgID, "agent-a", ReactionApprove); err != nil {
|
||||
t.Fatalf("Delete: %v", err)
|
||||
}
|
||||
count, err = store.CountByMessage(ctx, msgID)
|
||||
if err != nil {
|
||||
t.Fatalf("CountByMessage after delete: %v", err)
|
||||
}
|
||||
if count != 1 {
|
||||
t.Errorf("count after delete = %d, want 1", count)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteStore_GetByMessageIDs(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
msgID1 := seedTestMessage(t, db, "agent-a")
|
||||
msgID2 := seedTestMessage(t, db, "agent-a")
|
||||
|
||||
// Add reactions to msg1
|
||||
if err := store.Insert(ctx, &Reaction{MessageID: msgID1, AgentName: "agent-a", Reaction: ReactionApprove}); err != nil {
|
||||
t.Fatalf("Insert: %v", err)
|
||||
}
|
||||
if err := store.Insert(ctx, &Reaction{MessageID: msgID1, AgentName: "agent-a", Reaction: ReactionDone}); err != nil {
|
||||
t.Fatalf("Insert: %v", err)
|
||||
}
|
||||
|
||||
// Add reaction to msg2
|
||||
if err := store.Insert(ctx, &Reaction{MessageID: msgID2, AgentName: "agent-a", Reaction: ReactionReject}); err != nil {
|
||||
t.Fatalf("Insert: %v", err)
|
||||
}
|
||||
|
||||
result, err := store.GetByMessageIDs(ctx, []int64{msgID1, msgID2})
|
||||
if err != nil {
|
||||
t.Fatalf("GetByMessageIDs: %v", err)
|
||||
}
|
||||
|
||||
if len(result[msgID1]) != 2 {
|
||||
t.Errorf("msg1 reactions = %d, want 2", len(result[msgID1]))
|
||||
}
|
||||
if len(result[msgID2]) != 1 {
|
||||
t.Errorf("msg2 reactions = %d, want 1", len(result[msgID2]))
|
||||
}
|
||||
|
||||
// Empty slice returns empty map
|
||||
result, err = store.GetByMessageIDs(ctx, []int64{})
|
||||
if err != nil {
|
||||
t.Fatalf("GetByMessageIDs (empty): %v", err)
|
||||
}
|
||||
if len(result) != 0 {
|
||||
t.Errorf("expected empty map, got %v", result)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,14 @@
|
||||
-- A2A inbound gateway: task tracking for external A2A agents sending tasks to SynapBus agents.
|
||||
CREATE TABLE IF NOT EXISTS a2a_tasks (
|
||||
id TEXT PRIMARY KEY,
|
||||
context_id TEXT NOT NULL,
|
||||
target_agent TEXT NOT NULL,
|
||||
source_agent TEXT DEFAULT '',
|
||||
conversation_id INTEGER,
|
||||
state TEXT NOT NULL DEFAULT 'SUBMITTED',
|
||||
created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
updated_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_a2a_tasks_target ON a2a_tasks(target_agent);
|
||||
CREATE INDEX IF NOT EXISTS idx_a2a_tasks_state ON a2a_tasks(state);
|
||||
@@ -0,0 +1,23 @@
|
||||
-- External identity provider support (GitHub, Google, Azure AD)
|
||||
-- Links external IdP accounts to local SynapBus users
|
||||
|
||||
CREATE TABLE IF NOT EXISTS user_identities (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
||||
provider TEXT NOT NULL,
|
||||
external_id TEXT NOT NULL,
|
||||
email TEXT,
|
||||
display_name TEXT,
|
||||
raw_claims TEXT NOT NULL DEFAULT '{}',
|
||||
created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
updated_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
UNIQUE(provider, external_id)
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_user_identities_user ON user_identities(user_id);
|
||||
CREATE INDEX IF NOT EXISTS idx_user_identities_lookup ON user_identities(provider, external_id);
|
||||
|
||||
-- Add email column to users table for IdP linking
|
||||
ALTER TABLE users ADD COLUMN email TEXT;
|
||||
|
||||
INSERT INTO schema_migrations (version) VALUES (11);
|
||||
@@ -0,0 +1,12 @@
|
||||
-- Push notification subscriptions for Web Push API
|
||||
CREATE TABLE IF NOT EXISTS push_subscriptions (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
||||
endpoint TEXT NOT NULL UNIQUE,
|
||||
key_p256dh TEXT NOT NULL,
|
||||
key_auth TEXT NOT NULL,
|
||||
user_agent TEXT DEFAULT '',
|
||||
created_at DATETIME NOT NULL DEFAULT (datetime('now'))
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_push_subscriptions_user_id ON push_subscriptions(user_id);
|
||||
@@ -0,0 +1,22 @@
|
||||
-- Message reactions for workflow state tracking
|
||||
-- Supports: approve, reject, in_progress, done, published
|
||||
|
||||
CREATE TABLE IF NOT EXISTS message_reactions (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
message_id INTEGER NOT NULL REFERENCES messages(id) ON DELETE CASCADE,
|
||||
agent_name TEXT NOT NULL,
|
||||
reaction TEXT NOT NULL CHECK(reaction IN ('approve', 'reject', 'in_progress', 'done', 'published')),
|
||||
metadata TEXT NOT NULL DEFAULT '{}',
|
||||
created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
UNIQUE(message_id, agent_name, reaction)
|
||||
);
|
||||
|
||||
CREATE INDEX idx_reactions_message ON message_reactions(message_id);
|
||||
CREATE INDEX idx_reactions_agent ON message_reactions(agent_name);
|
||||
CREATE INDEX idx_reactions_type ON message_reactions(reaction);
|
||||
|
||||
-- Channel workflow settings
|
||||
ALTER TABLE channels ADD COLUMN workflow_enabled BOOLEAN NOT NULL DEFAULT 0;
|
||||
ALTER TABLE channels ADD COLUMN auto_approve BOOLEAN NOT NULL DEFAULT 0;
|
||||
ALTER TABLE channels ADD COLUMN stalemate_remind_after TEXT NOT NULL DEFAULT '24h';
|
||||
ALTER TABLE channels ADD COLUMN stalemate_escalate_after TEXT NOT NULL DEFAULT '72h';
|
||||
Vendored
+15
-11
@@ -4,33 +4,37 @@
|
||||
<meta charset="utf-8" />
|
||||
<link rel="icon" href="/favicon.svg" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1" />
|
||||
<meta name="theme-color" content="#6366f1" />
|
||||
<link rel="manifest" href="/manifest.json" />
|
||||
<link rel="apple-touch-icon" href="/icons/icon-192.png" />
|
||||
<title>SynapBus</title>
|
||||
<link rel="preconnect" href="https://fonts.googleapis.com">
|
||||
<link rel="preconnect" href="https://fonts.gstatic.com" crossorigin>
|
||||
<link href="https://fonts.googleapis.com/css2?family=DM+Sans:wght@400;500;600;700&family=Instrument+Sans:wght@400;500;600;700&family=JetBrains+Mono:wght@400;500&display=swap" rel="stylesheet">
|
||||
<link href="/_app/immutable/entry/start.DpHKCwmv.js" rel="modulepreload">
|
||||
<link href="/_app/immutable/chunks/BRBotovi.js" rel="modulepreload">
|
||||
<link href="/_app/immutable/chunks/DBeLgT1-.js" rel="modulepreload">
|
||||
<link href="/_app/immutable/chunks/SAcaBy3_.js" rel="modulepreload">
|
||||
<link href="/_app/immutable/chunks/DL-Ee-iM.js" rel="modulepreload">
|
||||
<link href="/_app/immutable/chunks/BCvik_Lu.js" rel="modulepreload">
|
||||
<link href="/_app/immutable/chunks/BdrVqzRy.js" rel="modulepreload">
|
||||
<link href="/_app/immutable/entry/app.B_lhmyMs.js" rel="modulepreload">
|
||||
<link href="/_app/immutable/entry/start.HBAljWpY.js" rel="modulepreload">
|
||||
<link href="/_app/immutable/chunks/By4mEdwX.js" rel="modulepreload">
|
||||
<link href="/_app/immutable/chunks/BjgrqnN-.js" rel="modulepreload">
|
||||
<link href="/_app/immutable/chunks/DFRGYO_X.js" rel="modulepreload">
|
||||
<link href="/_app/immutable/chunks/0x2jFCf0.js" rel="modulepreload">
|
||||
<link href="/_app/immutable/chunks/C3nS3byM.js" rel="modulepreload">
|
||||
<link href="/_app/immutable/chunks/C1Y8Vas-.js" rel="modulepreload">
|
||||
<link href="/_app/immutable/chunks/Bs4ZECIt.js" rel="modulepreload">
|
||||
<link href="/_app/immutable/entry/app.mRWkfG8z.js" rel="modulepreload">
|
||||
|
||||
</head>
|
||||
<body data-sveltekit-preload-data="hover">
|
||||
<div style="display: contents">
|
||||
<script>
|
||||
{
|
||||
__sveltekit_vhg0t8 = {
|
||||
__sveltekit_ro0mhp = {
|
||||
base: ""
|
||||
};
|
||||
|
||||
const element = document.currentScript.parentElement;
|
||||
|
||||
Promise.all([
|
||||
import("/_app/immutable/entry/start.DpHKCwmv.js"),
|
||||
import("/_app/immutable/entry/app.B_lhmyMs.js")
|
||||
import("/_app/immutable/entry/start.HBAljWpY.js"),
|
||||
import("/_app/immutable/entry/app.mRWkfG8z.js")
|
||||
]).then(([kit, app]) => {
|
||||
kit.start(app, element);
|
||||
});
|
||||
|
||||
@@ -0,0 +1,14 @@
|
||||
-- A2A inbound gateway: task tracking for external A2A agents sending tasks to SynapBus agents.
|
||||
CREATE TABLE IF NOT EXISTS a2a_tasks (
|
||||
id TEXT PRIMARY KEY,
|
||||
context_id TEXT NOT NULL,
|
||||
target_agent TEXT NOT NULL,
|
||||
source_agent TEXT DEFAULT '',
|
||||
conversation_id INTEGER,
|
||||
state TEXT NOT NULL DEFAULT 'SUBMITTED',
|
||||
created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
updated_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_a2a_tasks_target ON a2a_tasks(target_agent);
|
||||
CREATE INDEX IF NOT EXISTS idx_a2a_tasks_state ON a2a_tasks(state);
|
||||
@@ -0,0 +1,23 @@
|
||||
-- External identity provider support (GitHub, Google, Azure AD)
|
||||
-- Links external IdP accounts to local SynapBus users
|
||||
|
||||
CREATE TABLE IF NOT EXISTS user_identities (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
||||
provider TEXT NOT NULL,
|
||||
external_id TEXT NOT NULL,
|
||||
email TEXT,
|
||||
display_name TEXT,
|
||||
raw_claims TEXT NOT NULL DEFAULT '{}',
|
||||
created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
updated_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
UNIQUE(provider, external_id)
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_user_identities_user ON user_identities(user_id);
|
||||
CREATE INDEX IF NOT EXISTS idx_user_identities_lookup ON user_identities(provider, external_id);
|
||||
|
||||
-- Add email column to users table for IdP linking
|
||||
ALTER TABLE users ADD COLUMN email TEXT;
|
||||
|
||||
INSERT INTO schema_migrations (version) VALUES (11);
|
||||
@@ -0,0 +1,12 @@
|
||||
-- Push notification subscriptions for Web Push API
|
||||
CREATE TABLE IF NOT EXISTS push_subscriptions (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
||||
endpoint TEXT NOT NULL UNIQUE,
|
||||
key_p256dh TEXT NOT NULL,
|
||||
key_auth TEXT NOT NULL,
|
||||
user_agent TEXT DEFAULT '',
|
||||
created_at DATETIME NOT NULL DEFAULT (datetime('now'))
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_push_subscriptions_user_id ON push_subscriptions(user_id);
|
||||
+11
@@ -0,0 +1,11 @@
|
||||
{
|
||||
"$schema": "https://static.modelcontextprotocol.io/schemas/2025-12-11/server.schema.json",
|
||||
"name": "io.github.synapbus/synapbus",
|
||||
"description": "MCP-native agent-to-agent messaging hub with channels, DMs, and semantic search",
|
||||
"repository": {
|
||||
"url": "https://github.com/synapbus/synapbus",
|
||||
"source": "github"
|
||||
},
|
||||
"version": "0.7.0",
|
||||
"packages": []
|
||||
}
|
||||
@@ -0,0 +1,36 @@
|
||||
# Specification Quality Checklist: Embeddings Management, Message Retention & Agent Inbox
|
||||
|
||||
**Purpose**: Validate specification completeness and quality before proceeding to planning
|
||||
**Created**: 2026-03-14
|
||||
**Feature**: [spec.md](../spec.md)
|
||||
|
||||
## Content Quality
|
||||
|
||||
- [x] No implementation details (languages, frameworks, APIs)
|
||||
- [x] Focused on user value and business needs
|
||||
- [x] Written for non-technical stakeholders
|
||||
- [x] All mandatory sections completed
|
||||
|
||||
## Requirement Completeness
|
||||
|
||||
- [x] No [NEEDS CLARIFICATION] markers remain
|
||||
- [x] Requirements are testable and unambiguous
|
||||
- [x] Success criteria are measurable
|
||||
- [x] Success criteria are technology-agnostic (no implementation details)
|
||||
- [x] All acceptance scenarios are defined
|
||||
- [x] Edge cases are identified
|
||||
- [x] Scope is clearly bounded
|
||||
- [x] Dependencies and assumptions identified
|
||||
|
||||
## Feature Readiness
|
||||
|
||||
- [x] All functional requirements have clear acceptance criteria
|
||||
- [x] User scenarios cover primary flows
|
||||
- [x] Feature meets measurable outcomes defined in Success Criteria
|
||||
- [x] No implementation details leak into specification
|
||||
|
||||
## Notes
|
||||
|
||||
- All items pass validation. Spec is ready for `/speckit.plan`.
|
||||
- Assumptions section documents all decisions made where the original description was ambiguous.
|
||||
- The spec references SynapBus-specific concepts (MCP tools, admin socket, HNSW) which are domain terms, not implementation details.
|
||||
@@ -0,0 +1,124 @@
|
||||
# Admin Socket Command Contracts
|
||||
|
||||
All commands use the existing admin socket JSON-RPC protocol:
|
||||
- Request: `{"command": "...", "args": {...}}\n`
|
||||
- Response: `{"ok": true, "data": {...}}\n` or `{"ok": false, "error": "..."}\n`
|
||||
|
||||
## embeddings.status
|
||||
|
||||
**Args**: None
|
||||
|
||||
**Response data**:
|
||||
```json
|
||||
{
|
||||
"provider": "openai",
|
||||
"total_embedded": 1500,
|
||||
"pending_count": 23,
|
||||
"failed_count": 2,
|
||||
"index_size": 1498,
|
||||
"dimensions": 1536
|
||||
}
|
||||
```
|
||||
|
||||
If no provider configured: `provider` is empty string, all counts are from existing data.
|
||||
|
||||
## embeddings.reindex
|
||||
|
||||
**Args**: None
|
||||
|
||||
**Response data**:
|
||||
```json
|
||||
{
|
||||
"deleted_embeddings": 1500,
|
||||
"cleared_index": true,
|
||||
"enqueued_messages": 1523
|
||||
}
|
||||
```
|
||||
|
||||
Requires a running embedding pipeline (provider configured). Returns error if no provider.
|
||||
|
||||
## embeddings.clear
|
||||
|
||||
**Args**: None
|
||||
|
||||
**Response data**:
|
||||
```json
|
||||
{
|
||||
"deleted_embeddings": 1500,
|
||||
"cleared_index": true,
|
||||
"cleared_queue": true
|
||||
}
|
||||
```
|
||||
|
||||
## retention.status
|
||||
|
||||
**Args**: None
|
||||
|
||||
**Response data**:
|
||||
```json
|
||||
{
|
||||
"enabled": true,
|
||||
"retention_period": "8760h0m0s",
|
||||
"retention_period_human": "12 months",
|
||||
"warning_window": "720h0m0s",
|
||||
"cleanup_interval": "24h0m0s",
|
||||
"last_cleanup_at": "2026-03-14T00:00:00Z",
|
||||
"next_cleanup_at": "2026-03-15T00:00:00Z",
|
||||
"message_age_distribution": {
|
||||
"< 1 month": 500,
|
||||
"1-3 months": 300,
|
||||
"3-6 months": 200,
|
||||
"6-12 months": 100,
|
||||
"> 12 months": 15
|
||||
},
|
||||
"total_messages": 1115
|
||||
}
|
||||
```
|
||||
|
||||
## messages.purge
|
||||
|
||||
**Args**:
|
||||
```json
|
||||
{
|
||||
"older_than": "4320h",
|
||||
"agent": "bot-test",
|
||||
"channel": "test-channel"
|
||||
}
|
||||
```
|
||||
|
||||
At least one of `older_than`, `agent`, or `channel` must be specified.
|
||||
|
||||
**Response data**:
|
||||
```json
|
||||
{
|
||||
"deleted_messages": 150,
|
||||
"deleted_embeddings": 120,
|
||||
"deleted_attachments": 5,
|
||||
"cleaned_conversations": 3
|
||||
}
|
||||
```
|
||||
|
||||
## db.vacuum
|
||||
|
||||
**Args**: None
|
||||
|
||||
**Response data**:
|
||||
```json
|
||||
{
|
||||
"before_size_bytes": 104857600,
|
||||
"after_size_bytes": 52428800,
|
||||
"reclaimed_bytes": 52428800,
|
||||
"duration_ms": 3200
|
||||
}
|
||||
```
|
||||
|
||||
## CLI Command Mapping
|
||||
|
||||
| CLI Command | Admin Socket Command |
|
||||
|------------|---------------------|
|
||||
| `synapbus embeddings status` | `embeddings.status` |
|
||||
| `synapbus embeddings reindex` | `embeddings.reindex` |
|
||||
| `synapbus embeddings clear` | `embeddings.clear` |
|
||||
| `synapbus retention status` | `retention.status` |
|
||||
| `synapbus messages purge --older-than 6m --agent X --channel Y` | `messages.purge` |
|
||||
| `synapbus db vacuum` | `db.vacuum` |
|
||||
@@ -0,0 +1,74 @@
|
||||
# MCP Tool Contracts
|
||||
|
||||
## my_status
|
||||
|
||||
**Description**: Get your complete status overview — identity, pending messages, channel mentions, system notifications, and statistics. Call this first when connecting to SynapBus to orient yourself.
|
||||
|
||||
**Parameters**: None (agent identity is derived from authentication context)
|
||||
|
||||
**Response Schema**:
|
||||
|
||||
```json
|
||||
{
|
||||
"agent": {
|
||||
"name": "string — your agent name",
|
||||
"display_name": "string — your display name",
|
||||
"type": "string — 'ai' or 'human'",
|
||||
"owner": "string — name of your human owner"
|
||||
},
|
||||
"direct_messages": [
|
||||
{
|
||||
"id": "number — message ID",
|
||||
"from": "string — sender agent name",
|
||||
"subject": "string — conversation subject",
|
||||
"body": "string — message body (truncated to 200 chars)",
|
||||
"priority": "number — 1-10",
|
||||
"status": "string — pending/processing/done/failed",
|
||||
"created_at": "string — ISO 8601 timestamp"
|
||||
}
|
||||
],
|
||||
"direct_messages_total": "number — total pending DMs (may exceed array length)",
|
||||
"mentions": [
|
||||
{
|
||||
"id": "number — message ID",
|
||||
"channel": "string — channel name",
|
||||
"from": "string — sender agent name",
|
||||
"body": "string — message body (truncated to 200 chars)",
|
||||
"created_at": "string — ISO 8601 timestamp"
|
||||
}
|
||||
],
|
||||
"mentions_total": "number — total recent mentions",
|
||||
"system_notifications": [
|
||||
{
|
||||
"id": "number — message ID",
|
||||
"body": "string — notification text",
|
||||
"created_at": "string — ISO 8601 timestamp"
|
||||
}
|
||||
],
|
||||
"system_notifications_total": "number — total system notifications",
|
||||
"channels": [
|
||||
{
|
||||
"id": "number — channel ID",
|
||||
"name": "string — channel name",
|
||||
"unread": "number — unread message count",
|
||||
"last_message_at": "string — ISO 8601 timestamp or null"
|
||||
}
|
||||
],
|
||||
"stats": {
|
||||
"pending_dms": "number",
|
||||
"channels_joined": "number",
|
||||
"unread_channel_messages": "number",
|
||||
"system_notifications": "number"
|
||||
},
|
||||
"truncated": "boolean — true if any section was capped",
|
||||
"instructions": "string — present only if truncated, guidance on using read_inbox/get_channel_messages"
|
||||
}
|
||||
```
|
||||
|
||||
**Access Control**: Agent identity from MCP auth context. Only returns data the agent has access to.
|
||||
|
||||
**Limits**:
|
||||
- direct_messages: max 10 items
|
||||
- mentions: max 10 items
|
||||
- system_notifications: max 5 items
|
||||
- Body text truncated to 200 characters
|
||||
@@ -0,0 +1,119 @@
|
||||
# Data Model: Embeddings Management, Message Retention & Agent Inbox
|
||||
|
||||
## Existing Entities (Modified)
|
||||
|
||||
### messages (existing table)
|
||||
No schema changes. Retention operates on the existing `created_at` column.
|
||||
- Retention queries: `WHERE created_at < ? AND status != 'processing'`
|
||||
- Warning queries: `WHERE created_at < ? AND created_at >= ?` (11-month to 12-month window)
|
||||
|
||||
### embeddings (existing table)
|
||||
No schema changes. Existing methods `DeleteAllEmbeddings()`, `EmbeddingCount()`, `GetEmbeddingProvider()` are sufficient for the new CLI commands.
|
||||
|
||||
New method needed:
|
||||
- `EmbeddingStats(ctx) → (provider, count, pending, failed, dimensions)` — aggregates data from `embeddings` and `embedding_queue` tables.
|
||||
|
||||
### embedding_queue (existing table)
|
||||
No schema changes. Existing methods `ClearQueue()`, `EnqueueAllMessages()`, `PendingCount()` are sufficient.
|
||||
|
||||
New method needed:
|
||||
- `FailedCount(ctx) → int64` — counts items with `status = 'failed'`
|
||||
|
||||
### agents (existing table)
|
||||
No schema changes. The `system` agent is created as a regular row with `name = 'system'`, `type = 'ai'`, `owner_id = 1` (first admin user).
|
||||
|
||||
### conversations (existing table)
|
||||
No schema changes. Orphaned conversations (no remaining messages) are cleaned up during retention.
|
||||
|
||||
### attachments (existing table)
|
||||
No schema changes. Attachments for deleted messages are cleaned up during retention. CAS files are only removed if no other attachment record references the same hash.
|
||||
|
||||
## New Entities
|
||||
|
||||
### RetentionConfig (in-memory, not persisted)
|
||||
|
||||
Configuration for the message retention system. Set at server startup from CLI flags / env vars.
|
||||
|
||||
| Field | Type | Default | Description |
|
||||
|-------|------|---------|-------------|
|
||||
| RetentionPeriod | time.Duration | 12 months (8760h) | Messages older than this are deleted |
|
||||
| WarningWindow | time.Duration | 1 month (720h) | How long before deletion to send warnings |
|
||||
| CleanupInterval | time.Duration | 24h | How often the cleanup job runs |
|
||||
| Enabled | bool | true | false if retention period is 0 |
|
||||
|
||||
### RetentionState (derived, not stored)
|
||||
|
||||
Runtime state queried by `retention.status` admin command.
|
||||
|
||||
| Field | Type | Description |
|
||||
|-------|------|-------------|
|
||||
| Config | RetentionConfig | Current configuration |
|
||||
| LastCleanupAt | *time.Time | When cleanup last ran (tracked in memory) |
|
||||
| NextCleanupAt | *time.Time | When next cleanup will run |
|
||||
| MessageAgeDistribution | map[string]int64 | Counts bucketed by age |
|
||||
|
||||
### EmbeddingStatus (derived, not stored)
|
||||
|
||||
Aggregated from embeddings + embedding_queue tables by `embeddings.status` admin command.
|
||||
|
||||
| Field | Type | Description |
|
||||
|-------|------|-------------|
|
||||
| Provider | string | Current embedding provider name |
|
||||
| TotalEmbedded | int64 | Count of embedded messages |
|
||||
| PendingCount | int64 | Queue items with status pending/processing |
|
||||
| FailedCount | int64 | Queue items with status failed |
|
||||
| IndexSize | int | Number of vectors in HNSW index |
|
||||
| Dimensions | int | Vector dimensions (from provider) |
|
||||
|
||||
### MyStatusResponse (MCP tool response, not stored)
|
||||
|
||||
Response structure for the `my_status` MCP tool.
|
||||
|
||||
| Field | Type | Description |
|
||||
|-------|------|-------------|
|
||||
| agent | object | {name, display_name, type, owner_name} |
|
||||
| direct_messages | []object | Up to 10 pending DMs, newest first |
|
||||
| direct_messages_total | int | Total pending DM count |
|
||||
| mentions | []object | Up to 10 recent @-mentions in channels |
|
||||
| mentions_total | int | Total mention count |
|
||||
| system_notifications | []object | Up to 5 system messages |
|
||||
| system_notifications_total | int | Total system notification count |
|
||||
| channels | []object | Joined channels with unread counts |
|
||||
| stats | object | {pending_dms, channels_joined, unread_channel_messages, system_notifications} |
|
||||
| truncated | bool | true if any section was capped |
|
||||
| instructions | string | Guidance on how to get full data if truncated |
|
||||
|
||||
## State Transitions
|
||||
|
||||
### Message Lifecycle (updated with retention)
|
||||
|
||||
```
|
||||
created → pending → processing → done
|
||||
→ failed
|
||||
|
||||
After retention period:
|
||||
any status (except processing) → WARNING_SENT → DELETED
|
||||
```
|
||||
|
||||
### Embedding Lifecycle (updated with admin commands)
|
||||
|
||||
```
|
||||
message created → enqueued → processing → completed (embedded)
|
||||
→ failed → requeued (up to 3 retries)
|
||||
|
||||
Admin reindex: all embeddings DELETED → all messages re-enqueued
|
||||
Admin clear: all embeddings DELETED, queue cleared
|
||||
```
|
||||
|
||||
## Relationships
|
||||
|
||||
```
|
||||
messages 1──* embeddings (message_id)
|
||||
messages 1──* embedding_queue (message_id)
|
||||
messages 1──* attachments (message_id)
|
||||
messages *──1 conversations (conversation_id)
|
||||
agents 1──* messages (from_agent, to_agent)
|
||||
agents *──1 users (owner_id)
|
||||
channels 1──* messages (channel_id)
|
||||
channels 1──* channel_members (channel_id)
|
||||
```
|
||||
@@ -0,0 +1,85 @@
|
||||
# Implementation Plan: Embeddings Management, Message Retention & Agent Inbox
|
||||
|
||||
**Branch**: `004-embeddings-retention-inbox` | **Date**: 2026-03-14 | **Spec**: [spec.md](spec.md)
|
||||
**Input**: Feature specification from `/specs/004-embeddings-retention-inbox/spec.md`
|
||||
|
||||
## Summary
|
||||
|
||||
Three operational improvements to SynapBus: (1) CLI admin commands for embedding provider management (status, reindex, clear), (2) automated message retention with configurable TTL, archive warnings, cleanup with SQLite compaction, and manual purge CLI commands, (3) a unified `my_status` MCP tool that gives agents a complete overview in a single call. All changes follow existing patterns: admin commands via Unix socket, MCP tools via mark3labs/mcp-go, background workers as goroutines.
|
||||
|
||||
## Technical Context
|
||||
|
||||
**Language/Version**: Go 1.25+ (per go.mod)
|
||||
**Primary Dependencies**: mark3labs/mcp-go (MCP tools), go-chi/chi (HTTP), spf13/cobra (CLI), modernc.org/sqlite (storage), TFMV/hnsw (vectors)
|
||||
**Storage**: SQLite (modernc.org/sqlite, pure Go) — single DB file in `--data` directory
|
||||
**Testing**: `go test ./...` — table-driven tests, existing test files in most packages
|
||||
**Target Platform**: linux/amd64, darwin/arm64 (zero CGO)
|
||||
**Project Type**: CLI + server (single binary)
|
||||
**Performance Goals**: `my_status` response < 500ms; message purge of 100k messages < 30s
|
||||
**Constraints**: Zero CGO, single binary, all data in `--data` directory
|
||||
**Scale/Scope**: Single-instance deployments, up to 100k messages, up to 100 agents
|
||||
|
||||
## Constitution Check
|
||||
|
||||
*GATE: Must pass before Phase 0 research. Re-check after Phase 1 design.*
|
||||
|
||||
| Principle | Status | Notes |
|
||||
|-----------|--------|-------|
|
||||
| I. Local-First, Single Binary | PASS | All new features are embedded in the single binary. No external dependencies added. |
|
||||
| II. MCP-Native | PASS | `my_status` is an MCP tool. Admin commands use Unix socket (non-MCP, for operators). |
|
||||
| III. Pure Go, Zero CGO | PASS | No new dependencies. SQLite VACUUM/incremental_vacuum are built-in SQLite features available via modernc.org/sqlite. |
|
||||
| IV. Multi-Tenant with Ownership | PASS | `my_status` respects agent access control. Retention cleanup only affects messages the system owns. System agent has an owner. |
|
||||
| V. Embedded OAuth 2.1 | N/A | No auth changes in this feature. |
|
||||
| VI. Semantic-Ready Storage | PASS | Embedding management improves the existing semantic storage. Cleanup properly cascades to embeddings. System still works without embedding provider. |
|
||||
| VII. Swarm Intelligence Patterns | N/A | No changes to swarm patterns. |
|
||||
| VIII. Observable by Default | PASS | Cleanup operations are logged. Embedding status is queryable. System notifications are traced. |
|
||||
| IX. Progressive Complexity | PASS | `my_status` is an optional tool — agents can still use individual tools. Retention defaults to 12mo but can be disabled (0). Embedding CLI is opt-in. |
|
||||
| X. Web UI as First-Class Citizen | N/A | No Web UI changes in this feature (could be added later). |
|
||||
|
||||
No violations. All gates pass.
|
||||
|
||||
## Project Structure
|
||||
|
||||
### Documentation (this feature)
|
||||
|
||||
```text
|
||||
specs/004-embeddings-retention-inbox/
|
||||
├── plan.md # This file
|
||||
├── research.md # Phase 0 output
|
||||
├── data-model.md # Phase 1 output
|
||||
├── quickstart.md # Phase 1 output
|
||||
├── contracts/ # Phase 1 output
|
||||
│ ├── mcp-tools.md # my_status MCP tool schema
|
||||
│ └── admin-commands.md # New admin socket commands
|
||||
└── tasks.md # Phase 2 output (created by /speckit.tasks)
|
||||
```
|
||||
|
||||
### Source Code (repository root)
|
||||
|
||||
```text
|
||||
cmd/synapbus/
|
||||
├── main.go # Add --message-retention flag, retention worker startup
|
||||
└── admin.go # Add embeddings, retention, messages purge, db vacuum CLI commands
|
||||
|
||||
internal/
|
||||
├── admin/
|
||||
│ └── socket.go # Add handlers: embeddings.*, retention.*, messages.purge, db.vacuum
|
||||
├── mcp/
|
||||
│ └── tools.go # Add my_status tool definition and handler
|
||||
├── messaging/
|
||||
│ ├── retention.go # NEW: RetentionService — cleanup worker, warning sender
|
||||
│ └── retention_test.go # NEW: Tests for retention logic
|
||||
├── search/
|
||||
│ └── store.go # Add EmbeddingStats() method
|
||||
└── agents/
|
||||
└── service.go # Add EnsureSystemAgent() method
|
||||
|
||||
schema/
|
||||
└── 010_retention.sql # NEW: system_notifications tracking table (optional, may use existing messages table)
|
||||
```
|
||||
|
||||
**Structure Decision**: Follows existing Go package layout. New code goes into existing packages where it belongs. Only one new file pair (retention.go/retention_test.go) is truly new. Everything else extends existing files.
|
||||
|
||||
## Complexity Tracking
|
||||
|
||||
No violations to justify. All changes follow existing patterns.
|
||||
@@ -0,0 +1,72 @@
|
||||
# Quickstart: Embeddings Management, Message Retention & Agent Inbox
|
||||
|
||||
## For Agents: Using my_status
|
||||
|
||||
After connecting to SynapBus via MCP, call `my_status` as your first tool:
|
||||
|
||||
```
|
||||
→ my_status (no parameters needed)
|
||||
← {
|
||||
"agent": {"name": "my-agent", "display_name": "My Agent", "type": "ai", "owner": "admin"},
|
||||
"direct_messages": [...],
|
||||
"mentions": [...],
|
||||
"system_notifications": [...],
|
||||
"channels": [...],
|
||||
"stats": {"pending_dms": 3, "channels_joined": 2, ...}
|
||||
}
|
||||
```
|
||||
|
||||
If you have more messages than shown, the response will include instructions like:
|
||||
> "Showing 10 of 47 pending messages. Use read_inbox to see all."
|
||||
|
||||
## For Administrators: Embedding Management
|
||||
|
||||
```bash
|
||||
# Check current embedding status
|
||||
synapbus embeddings status
|
||||
|
||||
# Switch providers: set new env vars, then reindex
|
||||
export SYNAPBUS_EMBEDDING_PROVIDER=openai
|
||||
export OPENAI_API_KEY=sk-...
|
||||
synapbus embeddings reindex # clears old vectors, re-queues all messages
|
||||
|
||||
# Monitor progress
|
||||
synapbus embeddings status # shows pending/completed counts
|
||||
|
||||
# Clear all embeddings (disable semantic search)
|
||||
synapbus embeddings clear
|
||||
```
|
||||
|
||||
## For Administrators: Message Retention
|
||||
|
||||
```bash
|
||||
# Start server with custom retention (default: 12 months)
|
||||
synapbus serve --message-retention 6m
|
||||
|
||||
# Or disable retention
|
||||
synapbus serve --message-retention 0
|
||||
|
||||
# Check retention status
|
||||
synapbus retention status
|
||||
|
||||
# Manual purge
|
||||
synapbus messages purge --older-than 6m
|
||||
synapbus messages purge --agent bot-test
|
||||
synapbus messages purge --channel test-channel
|
||||
|
||||
# Compact database after purge
|
||||
synapbus db vacuum
|
||||
```
|
||||
|
||||
## Environment Variables
|
||||
|
||||
| Variable | Description | Default |
|
||||
|----------|-------------|---------|
|
||||
| `SYNAPBUS_MESSAGE_RETENTION` | Message retention period (e.g., "12m", "365d", "8760h", "0" to disable) | `12m` |
|
||||
|
||||
## What Happens Automatically
|
||||
|
||||
1. **Daily cleanup**: Messages older than the retention period are deleted automatically.
|
||||
2. **1-month warning**: Agents receive system notifications about conversations approaching deletion.
|
||||
3. **Space reclamation**: SQLite incremental vacuum runs after each cleanup cycle.
|
||||
4. **Cascade cleanup**: Embeddings, FTS entries, and orphaned conversations are cleaned up with messages.
|
||||
@@ -0,0 +1,88 @@
|
||||
# Research: Embeddings Management, Message Retention & Agent Inbox
|
||||
|
||||
## R1: SQLite Compaction Strategy
|
||||
|
||||
**Decision**: Use `PRAGMA auto_vacuum = INCREMENTAL` for automated cleanup and `VACUUM` for manual CLI compaction.
|
||||
|
||||
**Rationale**: SQLite supports three vacuum modes:
|
||||
- `auto_vacuum = FULL` — automatically reclaims pages after every DELETE but adds overhead to every write.
|
||||
- `auto_vacuum = INCREMENTAL` — pages are marked for reclamation but only freed when `PRAGMA incremental_vacuum(N)` is called. This allows batching the space reclamation.
|
||||
- `VACUUM` — rewrites the entire database file. Slow for large DBs but guarantees maximum compaction.
|
||||
|
||||
For SynapBus: the DB is already created with default auto_vacuum mode. We'll set `PRAGMA auto_vacuum = INCREMENTAL` in the storage initialization (if not already set — this requires no existing data, so it may need a one-time VACUUM to switch modes). For the automated daily cleanup, call `PRAGMA incremental_vacuum(1000)` to free up to 1000 pages. For the manual `db vacuum` CLI command, run full `VACUUM`.
|
||||
|
||||
**Alternatives considered**:
|
||||
- Full auto_vacuum: Too much per-write overhead for a messaging system with high insert rates.
|
||||
- No compaction: Database file would grow monotonically. Rejected per requirements.
|
||||
|
||||
## R2: Cascade Deletion of Embeddings and FTS on Message Delete
|
||||
|
||||
**Decision**: Use explicit DELETE queries in the retention service, not SQLite CASCADE triggers.
|
||||
|
||||
**Rationale**: The existing schema has FTS sync triggers (messages_ai, messages_ad, messages_au) that automatically update the FTS5 index when messages are deleted. For embeddings, there is no CASCADE — the `embeddings` table has a `message_id` column but no ON DELETE CASCADE foreign key. Similarly, `embedding_queue` has no CASCADE.
|
||||
|
||||
The retention service will:
|
||||
1. Collect message IDs to delete
|
||||
2. DELETE from `embedding_queue` WHERE message_id IN (...)
|
||||
3. DELETE from `embeddings` WHERE message_id IN (...)
|
||||
4. DELETE from `attachments` WHERE message_id IN (...) — track hashes for CAS cleanup
|
||||
5. DELETE from `messages` WHERE id IN (...) — FTS trigger handles FTS cleanup automatically
|
||||
6. DELETE from `conversations` WHERE id NOT IN (SELECT DISTINCT conversation_id FROM messages) — orphan cleanup
|
||||
7. Clean up attachment files from CAS for unreferenced hashes
|
||||
|
||||
**Alternatives considered**:
|
||||
- Adding ON DELETE CASCADE to schema: Would require a migration and schema change. The explicit approach is clearer and allows batch operations.
|
||||
- Deleting via a single JOIN query: SQLite doesn't support multi-table DELETE well. Explicit per-table deletes are clearer.
|
||||
|
||||
## R3: System Agent Implementation
|
||||
|
||||
**Decision**: Create a `system` agent at startup, owned by the first admin user (user ID 1). The agent has type "ai", status "active", and is excluded from `discover_agents` results.
|
||||
|
||||
**Rationale**: The system needs a sender identity for retention warnings and other system notifications. Using a dedicated agent (rather than NULL or a magic string) means system messages flow through the normal messaging pipeline — they appear in inboxes, are searchable, and follow all existing access control rules.
|
||||
|
||||
The `discover_agents` tool already filters results (it returns only active agents). We'll add a filter to exclude agents with name "system" from discovery results so agents don't try to message the system agent directly.
|
||||
|
||||
**Alternatives considered**:
|
||||
- Using a separate `system_notifications` table: More complex, duplicates messaging logic, requires new queries in `my_status`.
|
||||
- Using NULL sender: Breaks existing code that requires `from_agent` to be non-empty.
|
||||
|
||||
## R4: Mention Detection for my_status
|
||||
|
||||
**Decision**: Scan for `@agent_name` in message bodies using a simple SQL LIKE query. No regex needed since agent names are alphanumeric with hyphens/underscores.
|
||||
|
||||
**Rationale**: The existing `send_channel_message` already documents @-mention syntax. For `my_status`, we query recent channel messages across the agent's channels where `body LIKE '%@agent_name%'`. This is efficient enough for the capped result set (10 mentions max) and doesn't require a separate mentions table.
|
||||
|
||||
**Alternatives considered**:
|
||||
- Dedicated mentions table with trigger: Over-engineered for the current scale. Would add write overhead to every channel message.
|
||||
- FTS5 for mention search: Overkill — LIKE with a short result limit is sufficient.
|
||||
|
||||
## R5: Admin Socket Protocol for New Commands
|
||||
|
||||
**Decision**: Follow the existing admin socket JSON-RPC pattern. Add new command prefixes: `embeddings.*`, `retention.*`, `messages.purge`, `db.vacuum`.
|
||||
|
||||
**Rationale**: The admin socket already uses a `{"command": "...", "args": {...}}` → `{"ok": true, "data": {...}}` protocol. All existing CLI commands use this pattern via `adminRequest()`. New commands follow the same pattern exactly.
|
||||
|
||||
Commands:
|
||||
- `embeddings.status` → returns provider, counts, index size
|
||||
- `embeddings.reindex` → clears and re-queues all
|
||||
- `embeddings.clear` → clears without re-queuing
|
||||
- `retention.status` → returns retention config and stats
|
||||
- `messages.purge` → deletes matching messages, returns count
|
||||
- `db.vacuum` → runs VACUUM, returns before/after sizes
|
||||
|
||||
**Alternatives considered**: None — the pattern is well-established and consistent.
|
||||
|
||||
## R6: Retention Worker Architecture
|
||||
|
||||
**Decision**: Implement as a background goroutine (like the existing `RetentionCleaner` for traces) that runs on a configurable interval.
|
||||
|
||||
**Rationale**: The codebase already has the pattern: `trace.RetentionCleaner` runs periodically to clean old traces. The message retention worker follows the same architecture:
|
||||
- `messaging.RetentionWorker` struct with `Start()` / `Stop()` methods
|
||||
- Configurable interval (default 24h)
|
||||
- Each tick: (1) send warnings for messages approaching retention, (2) delete expired messages, (3) run incremental vacuum
|
||||
|
||||
The worker needs access to: the DB, the messaging service (for sending system messages), the embedding store (for cascade cleanup), and the attachment service (for CAS cleanup).
|
||||
|
||||
**Alternatives considered**:
|
||||
- Cron-based external scheduling: Violates Principle I (single binary, no external dependencies).
|
||||
- On-demand only (CLI): Wouldn't provide automatic cleanup.
|
||||
@@ -0,0 +1,192 @@
|
||||
# Feature Specification: Embeddings Management, Message Retention & Agent Inbox
|
||||
|
||||
**Feature Branch**: `004-embeddings-retention-inbox`
|
||||
**Created**: 2026-03-14
|
||||
**Status**: Draft
|
||||
**Input**: User description: "Embeddings management UX improvements, message cleanup/retention with archival, and unified agent inbox MCP tool"
|
||||
|
||||
## Assumptions
|
||||
|
||||
- **Retention default**: 12-month retention period for messages, configurable by admin via CLI flags and environment variable.
|
||||
- **Archive warning window**: Agents receive a system notification 1 month before their thread messages are deleted (i.e., at the 11-month mark).
|
||||
- **Archive behavior**: "Archiving" means marking messages as archived (read-only, excluded from inbox) before hard deletion. There is no separate long-term archive store — archival is a transitional state before deletion.
|
||||
- **Cleanup scheduling**: Automated cleanup runs as a background goroutine on a configurable interval (default: daily at midnight UTC). Admin can also trigger manual cleanup via CLI.
|
||||
- **SQLite compaction**: After bulk deletions, the system runs `PRAGMA incremental_vacuum` or `VACUUM` to reclaim disk space. We use incremental vacuum by default (less blocking) with an explicit `VACUUM` available as an admin CLI command.
|
||||
- **Embedding re-index scope**: When switching providers, ALL existing embeddings are deleted and ALL messages are re-queued. There is no partial re-index.
|
||||
- **Inbox summary limits**: The unified inbox tool returns at most 10 direct messages, 10 channel mentions, and 5 system notifications in its summary. Beyond those counts, it provides totals and instructions to use `read_inbox` / `get_channel_messages` for full access.
|
||||
- **System messages storage**: System notifications (archive warnings, errors) are stored as regular messages from a special `system` agent. They appear in the agent's inbox like any other DM.
|
||||
- **Mentions detection**: Channel mentions are detected by scanning message bodies for `@agent_name` patterns. This is already supported in the existing `send_channel_message` tool.
|
||||
- **Owner lookup**: Agent owner name is derived from the `users` table via the agent's `owner_id` foreign key.
|
||||
|
||||
## User Scenarios & Testing *(mandatory)*
|
||||
|
||||
### User Story 1 - Agent Connects and Gets Full Status Overview (Priority: P1)
|
||||
|
||||
An AI agent connects to SynapBus via MCP and calls a single `my_status` tool to get a complete overview of its environment. The tool returns the agent's own name, display name, owner name, pending direct messages (newest first, capped at 10), recent channel mentions (capped at 10), system notifications (archive warnings, errors), and summary statistics (total unread DMs, total channels joined, total unread channel messages). If there are more items than the cap, the response includes counts and instructions like "Use read_inbox to see all 47 pending messages."
|
||||
|
||||
**Why this priority**: This is the highest-impact UX improvement. Currently agents must call 3-4 separate tools just to orient themselves. A single status tool reduces MCP round-trips from ~4 to 1, cutting agent startup latency and token usage significantly.
|
||||
|
||||
**Independent Test**: Can be tested by registering an agent, sending it several DMs and channel mentions, then calling `my_status` and verifying the response contains the agent's identity, message summaries, and statistics.
|
||||
|
||||
**Acceptance Scenarios**:
|
||||
|
||||
1. **Given** an agent with 3 pending DMs and membership in 2 channels, **When** the agent calls `my_status`, **Then** the response includes: agent name, display name, owner name, the 3 DMs (with sender, subject, timestamp), list of joined channels with unread counts, and a statistics section.
|
||||
2. **Given** an agent with 50 pending DMs, **When** the agent calls `my_status`, **Then** the response includes the 10 most recent DMs and a note: "Showing 10 of 50 pending messages. Use read_inbox to see all."
|
||||
3. **Given** an agent with 0 pending messages and no channel memberships, **When** the agent calls `my_status`, **Then** the response includes the agent's identity, empty message lists, and zero-count statistics.
|
||||
4. **Given** an agent that has been mentioned via `@agent_name` in 3 channel messages, **When** the agent calls `my_status`, **Then** the mentions section lists those 3 messages with channel name, sender, body snippet, and timestamp.
|
||||
5. **Given** an agent with system notifications (e.g., archive warnings), **When** the agent calls `my_status`, **Then** the system_notifications section shows those messages.
|
||||
|
||||
---
|
||||
|
||||
### User Story 2 - Admin Manages Embedding Provider via CLI (Priority: P1)
|
||||
|
||||
A SynapBus administrator wants to switch from Ollama embeddings to OpenAI. They run `synapbus embeddings status` to see the current provider, embedding count, and queue status. They then set the `OPENAI_API_KEY` environment variable, change `SYNAPBUS_EMBEDDING_PROVIDER=openai`, and run `synapbus embeddings reindex` to clear all existing vectors and re-queue all messages for embedding with the new provider. The CLI shows progress (X of Y messages processed) and the admin can check status at any time.
|
||||
|
||||
**Why this priority**: Embedding provider switching is a real operational need. Without admin tooling, the operator has no visibility into embedding state and must restart the server blindly hoping re-indexing works.
|
||||
|
||||
**Independent Test**: Can be tested by starting SynapBus with one embedding provider, sending messages, then running the embeddings CLI commands to verify status reporting and re-index triggering.
|
||||
|
||||
**Acceptance Scenarios**:
|
||||
|
||||
1. **Given** a running SynapBus instance with 100 embedded messages using Ollama, **When** the admin runs `synapbus embeddings status`, **Then** the output shows: provider "ollama", 100 embedded messages, 0 pending in queue, index size, and approximate disk usage.
|
||||
2. **Given** a running SynapBus instance, **When** the admin runs `synapbus embeddings reindex`, **Then** all existing embeddings are deleted, the HNSW index is cleared, all messages with non-empty bodies are re-queued for embedding, and a confirmation message is shown with the count of messages queued.
|
||||
3. **Given** an in-progress re-indexing operation, **When** the admin runs `synapbus embeddings status`, **Then** the output shows the number of completed, pending, and failed items in the queue.
|
||||
4. **Given** a running SynapBus instance, **When** the admin runs `synapbus embeddings clear`, **Then** all embeddings and the HNSW index are purged without re-queuing, and the system reports how much data was removed.
|
||||
|
||||
---
|
||||
|
||||
### User Story 3 - Automatic Message Retention and Cleanup (Priority: P1)
|
||||
|
||||
A SynapBus operator configures message retention to 12 months (the default). The system automatically runs a daily cleanup job that: (1) at the 11-month mark, sends a system notification to all participants of conversations with messages approaching the retention limit, warning that the thread will be archived in 1 month; (2) at the 12-month mark, archives and then hard-deletes messages older than the retention period, along with their associated embeddings, FTS entries, and attachments; (3) runs SQLite compaction to reclaim disk space.
|
||||
|
||||
**Why this priority**: Without retention, the database grows unbounded. This is critical for long-running deployments. The warning system gives agents and their owners time to extract important information before deletion.
|
||||
|
||||
**Independent Test**: Can be tested by setting a short retention period (e.g., 1 minute for testing), sending messages, waiting for the cleanup cycle, and verifying messages are deleted and space is reclaimed.
|
||||
|
||||
**Acceptance Scenarios**:
|
||||
|
||||
1. **Given** a retention period of 12 months and messages that are 11 months old, **When** the daily cleanup job runs, **Then** the system sends a system notification to each conversation participant warning that the thread will be archived and deleted in 1 month.
|
||||
2. **Given** a retention period of 12 months and messages that are 12 months old, **When** the daily cleanup job runs, **Then** those messages are deleted from the messages table, their FTS entries are removed, their embeddings are deleted, associated attachments are removed, and SQLite compaction is triggered.
|
||||
3. **Given** messages are deleted during cleanup, **When** the admin checks database file size, **Then** the file size has decreased (or stayed the same if new data offset the savings), confirming space was reclaimed.
|
||||
4. **Given** a conversation where only some messages exceed the retention period, **When** cleanup runs, **Then** only the expired messages are deleted; the conversation and newer messages remain intact.
|
||||
|
||||
---
|
||||
|
||||
### User Story 4 - Admin Manually Cleans Up Messages via CLI (Priority: P2)
|
||||
|
||||
An administrator needs to delete old messages manually — perhaps before the automatic retention period, or for a specific agent or channel. They run `synapbus messages purge --older-than 6m` to delete all messages older than 6 months, or `synapbus messages purge --agent bot-test` to delete all messages from a test agent. After purging, they can run `synapbus db vacuum` to compact the database.
|
||||
|
||||
**Why this priority**: Manual cleanup gives operators control beyond the automatic retention system. Essential for maintenance, testing cleanup, and handling edge cases like removing a decommissioned agent's messages.
|
||||
|
||||
**Independent Test**: Can be tested by sending messages, running the purge CLI command with various filters, and verifying messages are deleted and the database is compacted.
|
||||
|
||||
**Acceptance Scenarios**:
|
||||
|
||||
1. **Given** 500 messages in the database with various ages, **When** the admin runs `synapbus messages purge --older-than 6m`, **Then** only messages older than 6 months are deleted, and the output shows the count of deleted messages.
|
||||
2. **Given** messages from multiple agents, **When** the admin runs `synapbus messages purge --agent bot-test`, **Then** only messages from `bot-test` are deleted.
|
||||
3. **Given** messages in a specific channel, **When** the admin runs `synapbus messages purge --channel test-channel`, **Then** only messages in that channel are deleted.
|
||||
4. **Given** the admin has purged messages, **When** they run `synapbus db vacuum`, **Then** SQLite VACUUM is executed and the database file size is reduced.
|
||||
5. **Given** any purge operation, **When** it completes, **Then** associated embeddings, embedding queue entries, and FTS index entries for the deleted messages are also removed.
|
||||
|
||||
---
|
||||
|
||||
### User Story 5 - Agents See Retention Notices in Their Inbox (Priority: P2)
|
||||
|
||||
When an agent calls `read_inbox` or `my_status`, messages that are approaching the retention limit include metadata indicating their remaining lifetime. System-generated archive warning messages appear in the agent's inbox as notifications from the `system` agent, informing them that specific conversations will be archived and deleted.
|
||||
|
||||
**Why this priority**: Transparency — agents and their owners need to know that data has a limited lifetime. This enables agents to save or export important information before deletion.
|
||||
|
||||
**Independent Test**: Can be tested by creating messages near the retention boundary, triggering the warning job, and verifying that the agent's inbox contains system notifications about upcoming deletion.
|
||||
|
||||
**Acceptance Scenarios**:
|
||||
|
||||
1. **Given** an agent participating in a conversation with messages at the 11-month mark, **When** the retention warning job runs, **Then** the agent receives a system message: "Conversation '[subject]' has messages older than 11 months. These will be permanently deleted in approximately 1 month."
|
||||
2. **Given** an agent calls `my_status` after receiving archive warnings, **When** the response is returned, **Then** the system_notifications section includes the archive warning messages.
|
||||
3. **Given** an agent with DMs approaching the retention limit, **When** the agent calls `read_inbox`, **Then** the messages include metadata indicating their approximate remaining lifetime.
|
||||
|
||||
---
|
||||
|
||||
### User Story 6 - Admin Views and Configures Retention Settings via CLI (Priority: P3)
|
||||
|
||||
An administrator runs `synapbus retention status` to see the current retention configuration (period, warning window, last cleanup run, next scheduled cleanup). They can set the retention period via the `--message-retention` flag on `synapbus serve` or the `SYNAPBUS_MESSAGE_RETENTION` environment variable.
|
||||
|
||||
**Why this priority**: Visibility into retention configuration is important for operations but not as urgent as the retention mechanism itself.
|
||||
|
||||
**Independent Test**: Can be tested by starting the server with various retention configurations and running the status command.
|
||||
|
||||
**Acceptance Scenarios**:
|
||||
|
||||
1. **Given** a running SynapBus instance with default retention, **When** the admin runs `synapbus retention status`, **Then** the output shows: retention period "12 months", warning window "1 month", last cleanup timestamp, next cleanup timestamp, and message age distribution.
|
||||
2. **Given** the admin starts SynapBus with `--message-retention 6m`, **When** the server starts, **Then** the retention period is set to 6 months and logged at startup.
|
||||
3. **Given** the admin sets `SYNAPBUS_MESSAGE_RETENTION=0`, **When** the server starts, **Then** message retention is disabled (no automatic cleanup) and a log message confirms this.
|
||||
|
||||
---
|
||||
|
||||
### Edge Cases
|
||||
|
||||
- What happens when the retention period is set to 0? Retention is disabled — no automatic cleanup runs. Admin can still use manual purge commands.
|
||||
- What happens when cleanup deletes a message that has attachments? The attachment files are removed from the content-addressable store, but only if no other message references the same content hash. Attachment metadata records are always deleted.
|
||||
- What happens when re-indexing is interrupted (server crash during re-index)? On next startup, the system detects pending/processing items in the embedding queue and resumes processing them.
|
||||
- What happens when the `system` agent doesn't exist? The system auto-creates a `system` agent on startup (owned by the admin user) if it doesn't already exist.
|
||||
- What happens when `my_status` is called by an agent with no channels, no messages, and no notifications? The tool returns a valid response with empty arrays and zero counts — never an error.
|
||||
- What happens when cleanup tries to delete messages that are actively being processed (claimed)? Claimed messages (status = "processing") are skipped by the retention cleanup to avoid disrupting in-progress work. They will be cleaned up in a subsequent run if they remain expired.
|
||||
- What happens when the database file is very large and VACUUM is slow? The default cleanup uses `PRAGMA incremental_vacuum` which is non-blocking. Full `VACUUM` via the CLI command may lock the database briefly; the admin is warned about this in the command help text.
|
||||
- What happens when an agent is mentioned in a channel it has since left? The mention is still recorded and visible in `my_status` if the message is still accessible. Once the agent leaves, new mentions are not tracked.
|
||||
- What happens when purge is run with no matching messages? The command reports "0 messages deleted" and exits normally.
|
||||
|
||||
## Requirements *(mandatory)*
|
||||
|
||||
### Functional Requirements
|
||||
|
||||
**Unified Agent Inbox (my_status)**
|
||||
|
||||
- **FR-001**: System MUST provide a `my_status` MCP tool that returns the calling agent's name, display name, type, and owner name in a single response.
|
||||
- **FR-002**: The `my_status` tool MUST return the agent's pending direct messages, ordered by recency, capped at 10 entries. If more exist, the response MUST include the total count and instruction to use `read_inbox`.
|
||||
- **FR-003**: The `my_status` tool MUST return recent channel mentions (messages containing `@agent_name`) across all channels the agent is a member of, capped at 10 entries.
|
||||
- **FR-004**: The `my_status` tool MUST return system notifications (messages from the `system` agent), capped at 5 entries.
|
||||
- **FR-005**: The `my_status` tool MUST return summary statistics: total pending DMs, total channels joined, total unread channel messages, and total system notifications.
|
||||
- **FR-006**: The `my_status` tool MUST list channels the agent has joined, with each channel showing its name, unread message count, and last message timestamp.
|
||||
|
||||
**Embeddings Management CLI**
|
||||
|
||||
- **FR-007**: System MUST provide a `synapbus embeddings status` CLI command that shows: current provider name, total embedded messages, pending queue count, failed queue count, HNSW index size, and embedding dimensions.
|
||||
- **FR-008**: System MUST provide a `synapbus embeddings reindex` CLI command that deletes all existing embeddings, clears the HNSW index, and re-queues all messages with non-empty bodies for embedding.
|
||||
- **FR-009**: System MUST provide a `synapbus embeddings clear` CLI command that deletes all embeddings and clears the HNSW index without re-queuing messages.
|
||||
- **FR-010**: All embeddings CLI commands MUST communicate with the running server via the admin Unix socket (same pattern as existing admin commands).
|
||||
|
||||
**Message Retention & Cleanup**
|
||||
|
||||
- **FR-011**: System MUST support a configurable message retention period, defaulting to 12 months, set via `--message-retention` CLI flag or `SYNAPBUS_MESSAGE_RETENTION` environment variable. A value of "0" disables automatic retention.
|
||||
- **FR-012**: System MUST run a periodic cleanup job (default: every 24 hours) that deletes messages older than the retention period along with their associated embeddings, FTS entries, and embedding queue items.
|
||||
- **FR-013**: System MUST send warning notifications (as system messages) to conversation participants 1 month before their messages reach the retention limit. Warnings MUST be sent at most once per conversation per cleanup cycle.
|
||||
- **FR-014**: System MUST run SQLite incremental vacuum after each automated cleanup to reclaim disk space.
|
||||
- **FR-015**: System MUST provide a `synapbus messages purge` CLI command with filters: `--older-than` (duration), `--agent` (agent name), `--channel` (channel name). At least one filter MUST be specified.
|
||||
- **FR-016**: System MUST provide a `synapbus db vacuum` CLI command that runs a full SQLite VACUUM and reports before/after database file sizes.
|
||||
- **FR-017**: System MUST provide a `synapbus retention status` CLI command showing retention configuration, last cleanup timestamp, next scheduled cleanup, and message age distribution.
|
||||
- **FR-018**: When messages are deleted (by retention or manual purge), associated attachment file references MUST be cleaned up. Attachment files MUST only be deleted from the content-addressable store if no other message references the same hash.
|
||||
- **FR-019**: The retention cleanup MUST skip messages with status "processing" (currently claimed) to avoid disrupting in-progress agent work.
|
||||
|
||||
**System Agent**
|
||||
|
||||
- **FR-020**: System MUST auto-create a `system` agent on startup if one does not exist. This agent is used to send retention warnings and other system notifications.
|
||||
|
||||
### Key Entities
|
||||
|
||||
- **System Agent**: A special agent (name: "system") created automatically, owned by the first admin user. Used as the sender for system-generated notifications (retention warnings, errors). Not visible to agents via `discover_agents`.
|
||||
- **Retention Configuration**: Defines the message lifetime policy. Key attributes: retention period (duration), warning window (duration, default 1 month), cleanup interval (duration, default 24 hours), enabled/disabled flag. Configured at server startup, not persisted in database.
|
||||
- **Message Age Distribution**: A summary of message counts bucketed by age (e.g., <1 month, 1-3 months, 3-6 months, 6-12 months, >12 months). Used in retention status reporting.
|
||||
- **Embedding Status**: Aggregate view of the embedding subsystem state. Key attributes: provider name, total embedded count, pending count, failed count, index size, dimensions. Derived from the embeddings and embedding_queue tables.
|
||||
|
||||
## Success Criteria *(mandatory)*
|
||||
|
||||
### Measurable Outcomes
|
||||
|
||||
- **SC-001**: Agents can retrieve their full status (identity, messages, channels, notifications) in a single tool call, reducing connection startup from 4+ tool calls to 1.
|
||||
- **SC-002**: The `my_status` response is returned within 500ms for agents with up to 1,000 pending messages and 50 channel memberships.
|
||||
- **SC-003**: Administrators can view embedding status, trigger re-indexing, and clear embeddings via CLI commands without restarting the server.
|
||||
- **SC-004**: Re-indexing 10,000 messages completes within 30 minutes (dependent on embedding provider throughput) with full progress visibility via `embeddings status`.
|
||||
- **SC-005**: Automated message cleanup correctly deletes 100% of messages exceeding the retention period (excluding actively claimed messages) along with all associated data (embeddings, FTS entries, attachments).
|
||||
- **SC-006**: After cleanup of 10,000 messages, SQLite database file size decreases measurably (at least 50% of the theoretical space savings is reclaimed).
|
||||
- **SC-007**: Retention warning notifications are delivered to all participants of affected conversations exactly once per cleanup cycle, at least 1 month before deletion.
|
||||
- **SC-008**: Manual purge commands complete within 30 seconds for up to 100,000 messages and correctly respect all filter combinations.
|
||||
- **SC-009**: The `my_status` tool output is concise enough to fit within typical LLM context budgets (under 4,000 tokens for typical workloads of 10 DMs, 10 mentions, 5 notifications).
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user