fix: restore verified unified kernel baseline
- fix securebus execute-once, redaction-safe persistence, and failure terminalization across main and subagent runs - harden session continuity, heartbeat isolation, and context-carried tool execution metadata without regressing eval workflows - restore full verification baseline with prompt/eval/tooling fixes and sync docs to shipped branch truth
This commit is contained in:
parent
1bdd82cc47
commit
161dd4e3ce
51 changed files with 2797 additions and 364 deletions
27
README.md
27
README.md
|
|
@ -83,9 +83,9 @@ flowchart TB
|
||||||
|
|
||||||
subgraph Security["Security"]
|
subgraph Security["Security"]
|
||||||
VLT[Vault XChaCha20] --> SS[SecretStore]
|
VLT[Vault XChaCha20] --> SS[SecretStore]
|
||||||
SS --> KR[Keyring / Env / File]
|
SS --> KR[Env key / planned keyring-file]
|
||||||
RED[Redactor] --> SB
|
RED[Redactor] --> SB
|
||||||
ZKP[Schnorr ZKP] -.-> SOCK[Daemon Socket]
|
ZKP[Planned ZKP] -.-> SOCK[Daemon Socket]
|
||||||
end
|
end
|
||||||
|
|
||||||
subgraph Bus["Message Bus"]
|
subgraph Bus["Message Bus"]
|
||||||
|
|
@ -118,16 +118,16 @@ flowchart TB
|
||||||
|
|
||||||
| Decision | Rationale |
|
| Decision | Rationale |
|
||||||
|----------|-----------|
|
|----------|-----------|
|
||||||
| **Isolated Tool Runtime** | All tool calls route through a `SecureBus` that enforces capability manifests, injects secrets, scans output for leaks, and writes audit logs. The LLM never sees raw secrets. Layer 1-2 are live today; daemon/WASM isolation remains optional roadmap work. See [ADR-001](docs/adr/001-isolated-tool-runtime.md). |
|
| **Isolated Tool Runtime** | All tool calls route through a `SecureBus` that mediates execution, applies recursion-depth policy, performs `arg:` secret injection, scans LLM-facing output for leaks, and writes audit logs. Broader network/filesystem capability enforcement plus daemon/WASM isolation remain follow-on work. See [ADR-001](docs/adr/001-isolated-tool-runtime.md). |
|
||||||
| **DAG executor** | LLMCompiler-style parallel tool dispatch. The planner builds a dependency DAG in a single inference pass; the executor dispatches independent nodes concurrently. Joiner synthesizes results. Replanning on failure. Falls back to ReAct for simple single-tool cases. |
|
| **DAG executor** | LLMCompiler-style parallel tool dispatch. The planner builds a dependency DAG in a single inference pass; the executor dispatches independent nodes concurrently. Joiner synthesizes results. Replanning on failure. Falls back to ReAct for simple single-tool cases. |
|
||||||
| **Vendored Fantasy SDK** | `charm.land/fantasy` vendored into `internal/fantasy/` via `go.mod` replace directive. Enables direct modification for streaming hooks, tool call repair, and progressive disclosure. |
|
| **Vendored Fantasy SDK** | `charm.land/fantasy` vendored into `internal/fantasy/` via `go.mod` replace directive. Enables direct modification for streaming hooks, tool call repair, and progressive disclosure. |
|
||||||
| **MemGPT + Projection Kernel** | Working context (hot), recall items (warm), archival chunks (cold, embedded + indexed), observational memory, immutable messages, DAG snapshots, runtime checkpoints, and an active-context projection builder that assembles the live turn context. Semantic ContextTree scoring and RLM reduction now participate in the hot path. |
|
| **MemGPT + Projection Kernel** | Working context (hot), recall items (warm), archival chunks (cold, embedded + indexed), observational memory, immutable messages, DAG snapshots, runtime checkpoints, and an active-context projection builder that assembles the live turn context. Semantic ContextTree scoring and RLM reduction now participate in the hot path. |
|
||||||
| **Progressive tool disclosure** | Agent sees only `tool_search` and `tool_call` meta-tools. Discovers actual tools on demand via fuzzy search. Cuts system prompt tokens for large registries. |
|
| **Progressive tool disclosure** | Gateway tools stay visible, and the agent gets a small query-aware direct toolset for the current request. Wider discovery still flows through `tool_search` / `tool_call` and dynamic promotion, which keeps prompt size down without hiding obvious direct actions. |
|
||||||
| **libSQL over modernc/sqlite** | Native F32_BLOB for vector storage, `libsql_vector_idx` for ANN search, FTS5 for full-text. Single database, no external vector DB dependency. |
|
| **libSQL over modernc/sqlite** | Native F32_BLOB for vector storage, `libsql_vector_idx` for ANN search, FTS5 for full-text. Single database, no external vector DB dependency. |
|
||||||
| **BLOB primary keys** | 16-byte UUIDv7 stored as BLOB. Compact, byte-comparable, monotonically sortable by creation time. |
|
| **BLOB primary keys** | 16-byte UUIDv7 stored as BLOB. Compact, byte-comparable, monotonically sortable by creation time. |
|
||||||
| **XChaCha20-Poly1305 vault** | Secrets encrypted at rest with AES-256-GCM or XChaCha20-Poly1305. Master key from OS keyring, env var, or file. Schnorr ZKP remains planned for daemon-mode authentication. |
|
| **XChaCha20-Poly1305 vault** | Secrets are encrypted at rest with XChaCha20-Poly1305. The current user-facing master-key flow is env-backed via `DRAGONSCALE_MASTER_KEY`; richer keyring/file-backed flows remain roadmap work. |
|
||||||
| **Goose migrations** | Schema managed by `pressly/goose/v3`. 17 versioned migrations currently cover core schema, FTS5, vector indexes, KV store, documents, audit log outcomes, conversations, runtime state, DAG/checkpoint data, map operators, memory edges, soft delete, immutable messages, and RL/task-completion tables. |
|
| **Goose migrations** | Schema managed by `pressly/goose/v3`. 17 versioned migrations currently cover core schema, FTS5, vector indexes, KV store, documents, audit log outcomes, conversations, runtime state, DAG/checkpoint data, map operators, memory edges, soft delete, immutable messages, and RL/task-completion tables. |
|
||||||
| **FlatBuffers command protocol** | Zero-copy serialized `ToolRequest`/`ToolResponse` for the ITR command vocabulary. Same binary format across in-process channels, Unix sockets (daemon mode), and wazero WASM host calls. |
|
| **FlatBuffers command protocol** | Zero-copy serialized `ToolRequest`/`ToolResponse` for the ITR command vocabulary. The binary format is live on the in-process command path today and is designed to extend to socket/WASM transports as those optional surfaces mature. |
|
||||||
|
|
||||||
## Project Layout
|
## Project Layout
|
||||||
|
|
||||||
|
|
@ -160,7 +160,7 @@ pkg/
|
||||||
├── memory/ # Memory system
|
├── memory/ # Memory system
|
||||||
│ ├── dag/ # DAG-based context budget compression
|
│ ├── dag/ # DAG-based context budget compression
|
||||||
│ ├── delegate/ # libSQL storage backend (FTS5, vector, capabilities)
|
│ ├── delegate/ # libSQL storage backend (FTS5, vector, capabilities)
|
||||||
│ ├── migrations/ # Goose versioned schema migrations (001–010)
|
│ ├── migrations/ # Goose versioned schema migrations (001–017)
|
||||||
│ ├── observation/ # Observational memory (observer, reflector, store)
|
│ ├── observation/ # Observational memory (observer, reflector, store)
|
||||||
│ ├── sqlc/ # sqlc config + generated code
|
│ ├── sqlc/ # sqlc config + generated code
|
||||||
│ └── store/ # MemoryStore, retrieval, chunking, scoring, queuing
|
│ └── store/ # MemoryStore, retrieval, chunking, scoring, queuing
|
||||||
|
|
@ -174,7 +174,7 @@ pkg/
|
||||||
├── tools/ # Tool registry, meta-tools, built-in tools, CapableTool
|
├── tools/ # Tool registry, meta-tools, built-in tools, CapableTool
|
||||||
├── voice/ # Groq Whisper voice transcription
|
├── voice/ # Groq Whisper voice transcription
|
||||||
└── worker/ # Background job worker
|
└── worker/ # Background job worker
|
||||||
skills/ # Built-in skills (weather, tmux, summarize, github, hardware)
|
cmd/dragonscale/workspace/skills/ # Embedded builtin skills packaged with the CLI
|
||||||
config/ # Example configuration files
|
config/ # Example configuration files
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|
@ -250,7 +250,8 @@ docker compose logs -f dragonscale-gateway
|
||||||
|
|
||||||
## Secret Management
|
## Secret Management
|
||||||
|
|
||||||
DragonScale encrypts secrets at rest with XChaCha20-Poly1305. The master key is sourced from an environment variable, OS keyring, or file.
|
DragonScale encrypts secrets at rest with XChaCha20-Poly1305 and stores them in `~/.dragonscale/secrets.json`.
|
||||||
|
Today the supported operator flow is env-backed: `dragonscale secret init` prints a hex key, and secret operations expect `DRAGONSCALE_MASTER_KEY` to be set.
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
dragonscale secret init # Generate a master key
|
dragonscale secret init # Generate a master key
|
||||||
|
|
@ -260,17 +261,17 @@ dragonscale secret delete <name> # Remove a secret
|
||||||
```
|
```
|
||||||
|
|
||||||
> [!WARNING]
|
> [!WARNING]
|
||||||
> This is a security-sensitive operation. The master key is used to encrypt and decrypt secrets. If you lose it, you will not be able to decrypt secrets.
|
> This is a security-sensitive operation. The master key is used to encrypt and decrypt secrets. If you lose it, you will not be able to decrypt existing secrets.
|
||||||
> You should should NEVER store the master key in a file or environment variable if possible.
|
> Treat the printed key like a root credential: do not commit it, paste it into logs, or leave it in shell history.
|
||||||
|
|
||||||
|
|
||||||
Set the master key: `export DRAGONSCALE_MASTER_KEY=<hex>`
|
Set the master key: `export DRAGONSCALE_MASTER_KEY=<hex>`
|
||||||
|
|
||||||
Tools declare which secrets they need via `CapableTool.Capabilities()`. The SecureBus injects secrets into tool execution context at runtime — the LLM never sees them. Tool output is scanned for leaked patterns before it reaches the agent loop.
|
Tools declare which secrets they need via `CapableTool.Capabilities()`. The SecureBus centrally supports `arg:` secret injection today, and it scans LLM-facing tool output for leaked patterns before results reach the agent loop. `env:` / `header:` injection modes remain tool-specific follow-up work.
|
||||||
|
|
||||||
## Daemon Mode
|
## Daemon Mode
|
||||||
|
|
||||||
For non-embedded deployments, the SecureBus can run in a separate privileged daemon process. The agent connects as an unprivileged client over a Unix domain socket.
|
DragonScale can start a standalone SecureBus daemon over a Unix domain socket for operator workflows, but the main `agent` / `gateway` runtime still uses in-process SecureBus today. Treat daemon mode as an optional deployment surface; socket-client integration and ZKP auth remain follow-on work.
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
dragonscale daemon start # Start daemon (foreground, Ctrl+C to stop)
|
dragonscale daemon start # Start daemon (foreground, Ctrl+C to stop)
|
||||||
|
|
|
||||||
14
ROADMAP.md
14
ROADMAP.md
|
|
@ -262,8 +262,8 @@ flowchart LR
|
||||||
- [ ] With this we can push all completed agent work to a "review" queue for human review and approval before being sent off
|
- [ ] With this we can push all completed agent work to a "review" queue for human review and approval before being sent off
|
||||||
- [ ] Example:
|
- [ ] Example:
|
||||||
- [ ] emails, sms, any work where we want to ensure consistency, order, etc.
|
- [ ] emails, sms, any work where we want to ensure consistency, order, etc.
|
||||||
- [ ] Add performance metrics to the eval harness
|
- [ ] Add richer performance metrics to the eval harness
|
||||||
- [ ] Raw metrics of tool calls, LLM calls, token counts, duration, etc
|
- [x] Raw metrics of tool calls, LLM calls, token counts, duration, etc
|
||||||
- [ ] Per-test scores
|
- [ ] Per-test scores
|
||||||
- [ ] Side-by-side comparison matrix
|
- [ ] Side-by-side comparison matrix
|
||||||
- [ ] Compare to other agent runtimes
|
- [ ] Compare to other agent runtimes
|
||||||
|
|
@ -287,9 +287,9 @@ flowchart LR
|
||||||
- [ ] spi
|
- [ ] spi
|
||||||
- [ ] pwm
|
- [ ] pwm
|
||||||
- [ ] etc
|
- [ ] etc
|
||||||
- [ ] Migrate to Cobra CLI framework
|
- [x] Migrate to Cobra CLI framework
|
||||||
- [ ] Use command-palette pattern for subcommands
|
- [x] Use command-palette pattern for subcommands
|
||||||
- [ ] keep cli commands as pure cli that calls into the application
|
- [x] keep cli commands as pure cli that calls into the application
|
||||||
- [ ] Migrate to errbuilder-go (ZanzyTHEbar)
|
- [ ] Migrate to errbuilder-go (ZanzyTHEbar)
|
||||||
- [ ] Migrate to assert-lib (ZanzyTHEbar)
|
- [ ] Migrate to assert-lib (ZanzyTHEbar)
|
||||||
- [ ] Implement SubAgent Profiles
|
- [ ] Implement SubAgent Profiles
|
||||||
|
|
@ -324,7 +324,9 @@ flowchart LR
|
||||||
- [x] RLM baseline: production context reduction over oversized projection segments
|
- [x] RLM baseline: production context reduction over oversized projection segments
|
||||||
- [ ] Full recursive DAG expansion and deeper memory-controller orchestration
|
- [ ] Full recursive DAG expansion and deeper memory-controller orchestration
|
||||||
- [ ] ReAct/DAG routing: automatic mode selection (`ModeReAct | ModeDAG | ModeAuto`)
|
- [ ] ReAct/DAG routing: automatic mode selection (`ModeReAct | ModeDAG | ModeAuto`)
|
||||||
- [ ] Layer 3: SecretStore + keyring-based secret management
|
- [~] Layer 3: SecretStore + keyring-based secret management
|
||||||
|
- [x] Env-backed encrypted secret store with XChaCha20-Poly1305 vault
|
||||||
|
- [ ] OS keyring / richer backend support
|
||||||
- [ ] Layer 4: Daemon mode + Schnorr ZKP authentication
|
- [ ] Layer 4: Daemon mode + Schnorr ZKP authentication
|
||||||
- [ ] Layer 5: wazero WASM isolates (pure Go, no CGO), `CodeExec` command variant
|
- [ ] Layer 5: wazero WASM isolates (pure Go, no CGO), `CodeExec` command variant
|
||||||
- [ ] Plug-in tool support: `pkg/tools/registry.go` — add `Search(query) []ToolInfo` for ToolSearch
|
- [ ] Plug-in tool support: `pkg/tools/registry.go` — add `Search(query) []ToolInfo` for ToolSearch
|
||||||
|
|
|
||||||
|
|
@ -14,7 +14,7 @@ This ADR mixes shipped kernel behavior with the longer-range secure execution ro
|
||||||
|
|
||||||
- Layer 1-2 SecureBus mediation is active for tool execution.
|
- Layer 1-2 SecureBus mediation is active for tool execution.
|
||||||
- FlatBuffers command vocabulary is live for the internal command surface.
|
- FlatBuffers command vocabulary is live for the internal command surface.
|
||||||
- Capability enforcement, secret injection, leak scanning, and audit logging are on the hot path.
|
- Recursion-depth policy validation, `arg:` secret injection, leak scanning, and audit logging are on the hot path.
|
||||||
- Dependency-aware parallel tool execution is active through the DAG tool runtime.
|
- Dependency-aware parallel tool execution is active through the DAG tool runtime.
|
||||||
- `pkg/rlm` is now wired into active-context assembly as a reducer for oversized DAG / recall / archival projection segments.
|
- `pkg/rlm` is now wired into active-context assembly as a reducer for oversized DAG / recall / archival projection segments.
|
||||||
|
|
||||||
|
|
@ -26,8 +26,7 @@ This ADR mixes shipped kernel behavior with the longer-range secure execution ro
|
||||||
|
|
||||||
## Context
|
## Context
|
||||||
|
|
||||||
DragonScale tools execute in-process with the agent loop. The `Vault` (XChaCha20-Poly1305) exists for encrypting secrets at rest, but there is no pipeline for injecting those secrets into tool execution. The `Redactor` scans for sensitive patterns, but only in log paths — not on tool output before it reaches the LLM. There is no privilege boundary
|
DragonScale tools execute in-process with the agent loop. The `Vault` (XChaCha20-Poly1305) encrypts secrets at rest, the SecureBus performs `arg:` secret injection for declared secret refs, and the redaction path scans LLM-facing tool output before it reaches the agent loop. The stronger privilege-boundary roadmap in this ADR still matters because broader network/filesystem policy enforcement, daemon separation, and WASM isolation are not yet fully shipped.
|
||||||
between the LLM-facing agent code and the tool execution path.
|
|
||||||
|
|
||||||
A compromised tool — via prompt injection, malicious skill, or supply chain attack — has the same memory-space access as the agent itself. This is the same class of vulnerability that led to the OpenClaw token exfiltration incident (Feb 2026), where malicious skills on ClawHub could read API keys from the host environment and exfiltrate them through tool output.
|
A compromised tool — via prompt injection, malicious skill, or supply chain attack — has the same memory-space access as the agent itself. This is the same class of vulnerability that led to the OpenClaw token exfiltration incident (Feb 2026), where malicious skills on ClawHub could read API keys from the host environment and exfiltrate them through tool output.
|
||||||
|
|
||||||
|
|
@ -217,34 +216,23 @@ This schema serves four purposes: (1) zero-copy reads eliminate serialization ov
|
||||||
|
|
||||||
### Layer 3: Secret Store + Keyring Integration
|
### Layer 3: Secret Store + Keyring Integration
|
||||||
|
|
||||||
The `SecretStore` maps logical secret names to encrypted ciphertext, persisted to a local file (`~/.dragonscale/secrets.enc`). The `Vault` handles encryption/decryption.
|
The `SecretStore` maps logical secret names to encrypted ciphertext, persisted today to `~/.dragonscale/secrets.json`. The `Vault` handles encryption/decryption.
|
||||||
|
|
||||||
The master key for the `Vault` is sourced from one of three backends, selected at
|
Current shipped master-key behavior is env-backed (`DRAGONSCALE_MASTER_KEY`) with an in-memory fallback for contexts that do not need persisted secret access. Richer OS-keyring, passphrase, or file-backed key management remains planned work.
|
||||||
onboarding:
|
|
||||||
|
|
||||||
| Backend | Platform | Mechanism |
|
**Current CLI surface**:
|
||||||
|---------------|--------------|----------------------------------------------|
|
|
||||||
| OS Keyring | Linux/macOS | libsecret (GNOME), kwallet (KDE), Keychain |
|
|
||||||
| Passphrase | Any | Argon2id KDF from user passphrase |
|
|
||||||
| File | Embedded | Raw key file with restricted permissions |
|
|
||||||
|
|
||||||
Keyring support is gated behind a build tag (`!embedded`) to avoid pulling in CGO or D-Bus dependencies on constrained platforms.
|
|
||||||
|
|
||||||
**CLI surface**:
|
|
||||||
|
|
||||||
```
|
```
|
||||||
dragonscale secret add <name> # interactive prompt for value
|
dragonscale secret add <name> # reads the value from stdin / prompt input
|
||||||
dragonscale secret list # names only, no values
|
dragonscale secret list # names only, no values
|
||||||
dragonscale secret delete <name>
|
dragonscale secret delete <name>
|
||||||
dragonscale secret export # encrypted backup
|
|
||||||
dragonscale secret import <file> # restore from backup
|
|
||||||
```
|
```
|
||||||
|
|
||||||
The `onboard` command is extended to include master key setup as part of the interactive wizard.
|
`secret export` / `secret import` and richer onboarding-backed key setup remain planned work.
|
||||||
|
|
||||||
### Layer 4: Daemon Mode + ZKP Authentication
|
### Layer 4: Daemon Mode + ZKP Authentication
|
||||||
|
|
||||||
For non-embedded deployments (desktop, server), the SecureBus can optionally run in a separate privileged daemon process. The agent loop connects as an unprivileged client over a Unix domain socket.
|
For non-embedded deployments (desktop, server), the SecureBus can optionally run in a separate daemon process over a Unix domain socket. Today that daemon surface exists as a standalone operational mode; the main `agent` / `gateway` runtime still executes against the in-process SecureBus path.
|
||||||
|
|
||||||
```
|
```
|
||||||
┌──────────────────┐ Unix Socket ┌──────────────────┐
|
┌──────────────────┐ Unix Socket ┌──────────────────┐
|
||||||
|
|
|
||||||
|
|
@ -33,14 +33,15 @@ type ResponseContent []Content
|
||||||
|
|
||||||
// Text returns the text content of the response.
|
// Text returns the text content of the response.
|
||||||
func (r ResponseContent) Text() string {
|
func (r ResponseContent) Text() string {
|
||||||
|
var builder strings.Builder
|
||||||
for _, c := range r {
|
for _, c := range r {
|
||||||
if c.GetType() == ContentTypeText {
|
if c.GetType() == ContentTypeText {
|
||||||
if textContent, ok := AsContentType[TextContent](c); ok {
|
if textContent, ok := AsContentType[TextContent](c); ok {
|
||||||
return textContent.Text
|
builder.WriteString(textContent.Text)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return ""
|
return builder.String()
|
||||||
}
|
}
|
||||||
|
|
||||||
// Reasoning returns all reasoning content parts.
|
// Reasoning returns all reasoning content parts.
|
||||||
|
|
|
||||||
17
internal/fantasy/model_test.go
Normal file
17
internal/fantasy/model_test.go
Normal file
|
|
@ -0,0 +1,17 @@
|
||||||
|
package fantasy
|
||||||
|
|
||||||
|
import "testing"
|
||||||
|
|
||||||
|
func TestResponseContentText_ConcatenatesMultipleTextParts(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
content := ResponseContent{
|
||||||
|
TextContent{Text: "first "},
|
||||||
|
ReasoningContent{Text: "internal only"},
|
||||||
|
TextContent{Text: "second"},
|
||||||
|
}
|
||||||
|
|
||||||
|
if got := content.Text(); got != "first second" {
|
||||||
|
t.Fatalf("expected concatenated text parts, got %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -281,8 +281,8 @@ func NewRegistry(_ string) []app.Task {
|
||||||
NewShellTask("fmt", "Format Go code", staticGoScript("fmt ./..."), nil),
|
NewShellTask("fmt", "Format Go code", staticGoScript("fmt ./..."), nil),
|
||||||
NewCommandTask("lint", "Run all linting checks", lintSpecs, nil, nil),
|
NewCommandTask("lint", "Run all linting checks", lintSpecs, nil, nil),
|
||||||
NewShellTask("hooks", "Install git hooks", hooksScript, nil),
|
NewShellTask("hooks", "Install git hooks", hooksScript, nil),
|
||||||
NewShellTask("deps", "Download dependencies", staticGoScript("mod download && mod verify"), nil),
|
NewShellTask("deps", "Download dependencies", staticGoScript("mod download && $GO mod verify"), nil),
|
||||||
NewShellTask("update-deps", "Update dependencies", staticGoScript("get -u ./... && mod tidy"), nil),
|
NewShellTask("update-deps", "Update dependencies", staticGoScript("get -u ./... && $GO mod tidy"), nil),
|
||||||
NewShellTask("sqlc-check", "Verify sqlc generation is idempotent", sqlcCheckScript, nil),
|
NewShellTask("sqlc-check", "Verify sqlc generation is idempotent", sqlcCheckScript, nil),
|
||||||
NewShellTask("flatc-check", "Verify flatc generation is idempotent", flatcCheckScript, nil),
|
NewShellTask("flatc-check", "Verify flatc generation is idempotent", flatcCheckScript, nil),
|
||||||
NewShellTask("sqlc-vet", "Run sqlc vet rules", sqlcVetScript, nil),
|
NewShellTask("sqlc-vet", "Run sqlc vet rules", sqlcVetScript, nil),
|
||||||
|
|
@ -734,6 +734,9 @@ func evalRunSpecs(c *app.Context) []runner.CommandSpec {
|
||||||
}
|
}
|
||||||
|
|
||||||
specs := append([]runner.CommandSpec{}, maybeEvalBuildSpecs(c)...)
|
specs := append([]runner.CommandSpec{}, maybeEvalBuildSpecs(c)...)
|
||||||
|
if hasEvalSourceTree(c) {
|
||||||
|
specs = append(specs, evalFixturesSpecs(c)...)
|
||||||
|
}
|
||||||
if debug && strings.TrimSpace(baseCfg) != "" {
|
if debug && strings.TrimSpace(baseCfg) != "" {
|
||||||
specs = append(specs, runner.CommandSpec{
|
specs = append(specs, runner.CommandSpec{
|
||||||
Name: "echo",
|
Name: "echo",
|
||||||
|
|
|
||||||
|
|
@ -191,6 +191,64 @@ func TestBuildAllTaskRejectsNonLinuxTarget(t *testing.T) {
|
||||||
require.Len(t, fake.Calls, 0)
|
require.Len(t, fake.Calls, 0)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestDepsTaskPrefixesEachGoCommand(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
ctx := &app.Context{
|
||||||
|
Root: t.TempDir(),
|
||||||
|
ExtraEnv: map[string]string{
|
||||||
|
"SKIP_DEVCONTAINER_WRAPPER": "1",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
fake := &runner.FakeRunner{Result: runner.CommandResult{ExitCode: 0}}
|
||||||
|
|
||||||
|
var deps app.Task
|
||||||
|
for _, task := range NewRegistry(ctx.Root) {
|
||||||
|
if task.Name() == "deps" {
|
||||||
|
deps = task
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
require.NotNil(t, deps)
|
||||||
|
|
||||||
|
_, err := deps.Run(context.Background(), fake, ctx)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, fake.Calls, 1)
|
||||||
|
|
||||||
|
script := strings.Join(fake.Calls[0].Args, " ")
|
||||||
|
require.Contains(t, script, "$GO mod download && $GO mod verify")
|
||||||
|
require.NotContains(t, script, "$GO mod download && mod verify")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUpdateDepsTaskPrefixesEachGoCommand(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
ctx := &app.Context{
|
||||||
|
Root: t.TempDir(),
|
||||||
|
ExtraEnv: map[string]string{
|
||||||
|
"SKIP_DEVCONTAINER_WRAPPER": "1",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
fake := &runner.FakeRunner{Result: runner.CommandResult{ExitCode: 0}}
|
||||||
|
|
||||||
|
var updateDeps app.Task
|
||||||
|
for _, task := range NewRegistry(ctx.Root) {
|
||||||
|
if task.Name() == "update-deps" {
|
||||||
|
updateDeps = task
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
require.NotNil(t, updateDeps)
|
||||||
|
|
||||||
|
_, err := updateDeps.Run(context.Background(), fake, ctx)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, fake.Calls, 1)
|
||||||
|
|
||||||
|
script := strings.Join(fake.Calls[0].Args, " ")
|
||||||
|
require.Contains(t, script, "$GO get -u ./... && $GO mod tidy")
|
||||||
|
require.NotContains(t, script, "$GO get -u ./... && mod tidy")
|
||||||
|
}
|
||||||
|
|
||||||
func TestEvalRunSpecsPreservesEvalConfig(t *testing.T) {
|
func TestEvalRunSpecsPreservesEvalConfig(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
|
|
@ -461,12 +519,17 @@ func TestEvalRunSpecsPrependsBuildWhenRunnerMissingAndSourceTreeExists(t *testin
|
||||||
root := t.TempDir()
|
root := t.TempDir()
|
||||||
require.NoError(t, os.MkdirAll(filepath.Join(root, "eval", "cmd", "eval-runner"), 0o755))
|
require.NoError(t, os.MkdirAll(filepath.Join(root, "eval", "cmd", "eval-runner"), 0o755))
|
||||||
require.NoError(t, os.WriteFile(filepath.Join(root, "eval", "cmd", "eval-runner", "main.go"), []byte("package main\nfunc main() {}\n"), 0o644))
|
require.NoError(t, os.WriteFile(filepath.Join(root, "eval", "cmd", "eval-runner", "main.go"), []byte("package main\nfunc main() {}\n"), 0o644))
|
||||||
|
require.NoError(t, os.MkdirAll(filepath.Join(root, "eval", "fixtures", "skills"), 0o755))
|
||||||
|
require.NoError(t, os.WriteFile(filepath.Join(root, "eval", "fixtures", "sample_data.txt"), []byte("fixture"), 0o644))
|
||||||
|
|
||||||
specs := evalRunSpecs(&app.Context{Root: root})
|
specs := evalRunSpecs(&app.Context{Root: root})
|
||||||
require.GreaterOrEqual(t, len(specs), 4)
|
require.GreaterOrEqual(t, len(specs), 10)
|
||||||
require.Equal(t, "go", specs[0].Name)
|
require.Equal(t, "go", specs[0].Name)
|
||||||
require.Equal(t, []string{"generate", "./..."}, specs[0].Args)
|
require.Equal(t, []string{"generate", "./..."}, specs[0].Args)
|
||||||
require.Equal(t, root, specs[0].Dir)
|
require.Equal(t, root, specs[0].Dir)
|
||||||
|
require.Equal(t, "mkdir", specs[3].Name)
|
||||||
|
require.Equal(t, "rm", specs[4].Name)
|
||||||
|
require.Equal(t, "cp", specs[8].Name)
|
||||||
require.Equal(t, "npx", specs[len(specs)-1].Name)
|
require.Equal(t, "npx", specs[len(specs)-1].Name)
|
||||||
require.Equal(t, filepath.Join(root, "eval"), specs[len(specs)-1].Dir)
|
require.Equal(t, filepath.Join(root, "eval"), specs[len(specs)-1].Dir)
|
||||||
}
|
}
|
||||||
|
|
@ -479,15 +542,39 @@ func TestEvalRunSpecsPrependsBuildWhenRunnerAlreadyExists(t *testing.T) {
|
||||||
require.NoError(t, os.WriteFile(filepath.Join(root, "eval", "cmd", "eval-runner", "main.go"), []byte("package main\nfunc main() {}\n"), 0o644))
|
require.NoError(t, os.WriteFile(filepath.Join(root, "eval", "cmd", "eval-runner", "main.go"), []byte("package main\nfunc main() {}\n"), 0o644))
|
||||||
require.NoError(t, os.MkdirAll(filepath.Join(root, "eval", "bin"), 0o755))
|
require.NoError(t, os.MkdirAll(filepath.Join(root, "eval", "bin"), 0o755))
|
||||||
require.NoError(t, os.WriteFile(filepath.Join(root, "eval", "bin", "eval-runner"), []byte("stale-binary"), 0o755))
|
require.NoError(t, os.WriteFile(filepath.Join(root, "eval", "bin", "eval-runner"), []byte("stale-binary"), 0o755))
|
||||||
|
require.NoError(t, os.MkdirAll(filepath.Join(root, "eval", "fixtures", "skills"), 0o755))
|
||||||
|
require.NoError(t, os.WriteFile(filepath.Join(root, "eval", "fixtures", "sample_data.txt"), []byte("fixture"), 0o644))
|
||||||
|
|
||||||
specs := evalRunSpecs(&app.Context{Root: root})
|
specs := evalRunSpecs(&app.Context{Root: root})
|
||||||
require.GreaterOrEqual(t, len(specs), 4)
|
require.GreaterOrEqual(t, len(specs), 10)
|
||||||
require.Equal(t, "go", specs[0].Name)
|
require.Equal(t, "go", specs[0].Name)
|
||||||
require.Equal(t, []string{"generate", "./..."}, specs[0].Args)
|
require.Equal(t, []string{"generate", "./..."}, specs[0].Args)
|
||||||
require.Equal(t, root, specs[0].Dir)
|
require.Equal(t, root, specs[0].Dir)
|
||||||
require.Equal(t, "npx", specs[len(specs)-1].Name)
|
require.Equal(t, "npx", specs[len(specs)-1].Name)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestEvalRunSpecsPrependsFixturesWhenSourceTreeExists(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
root := t.TempDir()
|
||||||
|
require.NoError(t, os.MkdirAll(filepath.Join(root, "eval", "cmd", "eval-runner"), 0o755))
|
||||||
|
require.NoError(t, os.WriteFile(filepath.Join(root, "eval", "cmd", "eval-runner", "main.go"), []byte("package main\nfunc main() {}\n"), 0o644))
|
||||||
|
require.NoError(t, os.MkdirAll(filepath.Join(root, "eval", "fixtures", "skills"), 0o755))
|
||||||
|
require.NoError(t, os.WriteFile(filepath.Join(root, "eval", "fixtures", "sample_data.txt"), []byte("fixture"), 0o644))
|
||||||
|
|
||||||
|
specs := evalRunSpecs(&app.Context{Root: root})
|
||||||
|
require.GreaterOrEqual(t, len(specs), 10)
|
||||||
|
require.Equal(t, "mkdir", specs[3].Name)
|
||||||
|
require.Equal(t, []string{"-p", filepath.Join(homeDir(), ".local", "share", "dragonscale", "sandbox")}, specs[3].Args)
|
||||||
|
require.Equal(t, "rm", specs[4].Name)
|
||||||
|
require.Equal(t, "bash", specs[7].Name)
|
||||||
|
require.Equal(t, "cp", specs[8].Name)
|
||||||
|
require.Equal(t, "bash", specs[9].Name)
|
||||||
|
require.Equal(t, "npx", specs[len(specs)-1].Name)
|
||||||
|
joinedArgs := strings.Join(specs[len(specs)-1].Args, " ")
|
||||||
|
require.Contains(t, joinedArgs, "promptfoo eval --config promptfooconfig.yaml")
|
||||||
|
}
|
||||||
|
|
||||||
func TestEvalTasksUseRepoRootPathsWhenCwdIsNested(t *testing.T) {
|
func TestEvalTasksUseRepoRootPathsWhenCwdIsNested(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -122,7 +122,7 @@ func (b *DefaultActiveContextBuilder) BuildTurnContext(ctx context.Context, req
|
||||||
summary = strings.TrimSpace(b.sessions.GetSummary(req.ProjectionRequest.SessionKey))
|
summary = strings.TrimSpace(b.sessions.GetSummary(req.ProjectionRequest.SessionKey))
|
||||||
}
|
}
|
||||||
|
|
||||||
systemSegments := b.buildSystemSegments(req.ProjectionRequest.SessionKey, summary, budget.System)
|
systemSegments := b.buildSystemSegments(req.ProjectionRequest.SessionKey, req.CurrentMessage, summary, budget.System)
|
||||||
projection.Segments = append(projection.Segments, systemSegments...)
|
projection.Segments = append(projection.Segments, systemSegments...)
|
||||||
|
|
||||||
var immutableHistory []*memory.ImmutableMessage
|
var immutableHistory []*memory.ImmutableMessage
|
||||||
|
|
@ -153,14 +153,14 @@ func (b *DefaultActiveContextBuilder) BuildTurnContext(ctx context.Context, req
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *DefaultActiveContextBuilder) buildSystemSegments(sessionKey, summary string, budget int) []memory.ProjectionSegment {
|
func (b *DefaultActiveContextBuilder) buildSystemSegments(sessionKey, currentMessage, summary string, budget int) []memory.ProjectionSegment {
|
||||||
if b.contextBuilder == nil || budget <= 0 {
|
if b.contextBuilder == nil || budget <= 0 {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
candidates := make([]memory.ProjectionSegment, 0, 2)
|
candidates := make([]memory.ProjectionSegment, 0, 2)
|
||||||
|
|
||||||
systemPrompt := strings.TrimSpace(b.contextBuilder.BuildSystemPromptWithBudget(0))
|
systemPrompt := strings.TrimSpace(b.contextBuilder.BuildSystemPromptForTurn(sessionKey, currentMessage, 0))
|
||||||
if systemPrompt != "" {
|
if systemPrompt != "" {
|
||||||
candidates = append(candidates, memory.ProjectionSegment{
|
candidates = append(candidates, memory.ProjectionSegment{
|
||||||
Kind: memory.ProjectionSegmentSystem,
|
Kind: memory.ProjectionSegmentSystem,
|
||||||
|
|
|
||||||
|
|
@ -69,7 +69,9 @@ func TestAssembleContext_UsesActiveContextProjection(t *testing.T) {
|
||||||
assert.Contains(t, projectionKinds(ac.projection), memory.ProjectionSegmentSystem)
|
assert.Contains(t, projectionKinds(ac.projection), memory.ProjectionSegmentSystem)
|
||||||
assert.Contains(t, projectionKinds(ac.projection), memory.ProjectionSegmentRecent)
|
assert.Contains(t, projectionKinds(ac.projection), memory.ProjectionSegmentRecent)
|
||||||
assert.Contains(t, projectionKinds(ac.projection), memory.ProjectionSegmentTool)
|
assert.Contains(t, projectionKinds(ac.projection), memory.ProjectionSegmentTool)
|
||||||
assert.Contains(t, ac.systemPrompt, "You have access to the following tools")
|
assert.Contains(t, ac.systemPrompt, "Plans vs actions")
|
||||||
|
assert.Contains(t, ac.systemPrompt, "Direct tool routing")
|
||||||
|
assert.NotContains(t, ac.systemPrompt, "You have access to the following tools")
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestActiveContextBuilder_IncludesPersistedDAGProjection(t *testing.T) {
|
func TestActiveContextBuilder_IncludesPersistedDAGProjection(t *testing.T) {
|
||||||
|
|
@ -150,6 +152,44 @@ func TestActiveContextBuilder_IncludesPersistedDAGProjection(t *testing.T) {
|
||||||
assert.Contains(t, projectionKinds(built.Projection), memory.ProjectionSegmentDAG)
|
assert.Contains(t, projectionKinds(built.Projection), memory.ProjectionSegmentDAG)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestActiveContextBuilder_UsesTurnSpecificToolHintsInSystemSegment(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
tmpDir, err := os.MkdirTemp("", "active-context-tools-*")
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer os.RemoveAll(tmpDir)
|
||||||
|
|
||||||
|
cfg := config.DefaultConfig()
|
||||||
|
cfg.Agents.Defaults.Sandbox = tmpDir
|
||||||
|
cfg.Agents.Defaults.MaxTokens = 4096
|
||||||
|
cfg.Memory.DBPath = filepath.Join(tmpDir, "active-context-tools.db")
|
||||||
|
|
||||||
|
al := mustNewAgentLoop(t, cfg, bus.NewMessageBus(), newMockLanguageModel("ok"))
|
||||||
|
require.NotNil(t, al.activeContextBuilder)
|
||||||
|
|
||||||
|
built, err := al.activeContextBuilder.BuildTurnContext(t.Context(), TurnContextBuildRequest{
|
||||||
|
ProjectionRequest: memory.ProjectionRequest{
|
||||||
|
AgentID: pkgroot.NAME,
|
||||||
|
SessionKey: "tool-hint-session",
|
||||||
|
MaxTokens: 4096,
|
||||||
|
},
|
||||||
|
CurrentMessage: "Capture these commitments and give me a reminder/follow-up plan with explicit timing.",
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, built.Projection)
|
||||||
|
|
||||||
|
var systemText string
|
||||||
|
for _, seg := range built.Projection.Segments {
|
||||||
|
if seg.Kind == memory.ProjectionSegmentSystem && seg.Source == "runtime_system" {
|
||||||
|
systemText = seg.Text
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
require.NotEmpty(t, systemText)
|
||||||
|
assert.Contains(t, systemText, "`memory`")
|
||||||
|
assert.NotContains(t, systemText, "`obligation`")
|
||||||
|
}
|
||||||
|
|
||||||
func insertImmutableMessage(t *testing.T, al *AgentLoop, msg *memory.ImmutableMessage) {
|
func insertImmutableMessage(t *testing.T, al *AgentLoop, msg *memory.ImmutableMessage) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
require.NoError(t, al.memDelegate.InsertImmutableMessage(t.Context(), msg))
|
require.NoError(t, al.memDelegate.InsertImmutableMessage(t.Context(), msg))
|
||||||
|
|
|
||||||
|
|
@ -4,6 +4,7 @@ package agent
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"database/sql"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
|
@ -26,6 +27,34 @@ import (
|
||||||
"golang.org/x/sync/errgroup"
|
"golang.org/x/sync/errgroup"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
const conversationBindingKVPrefix = "session:conversation:"
|
||||||
|
|
||||||
|
func conversationBindingKey(sessionKey string) string {
|
||||||
|
return conversationBindingKVPrefix + sessionKey
|
||||||
|
}
|
||||||
|
|
||||||
|
func (al *AgentLoop) loadBoundConversationID(ctx context.Context, sessionKey string) (ids.UUID, error) {
|
||||||
|
if al == nil || al.memDelegate == nil {
|
||||||
|
return ids.UUID{}, nil
|
||||||
|
}
|
||||||
|
raw, err := al.memDelegate.GetKV(ctx, pkg.NAME, conversationBindingKey(sessionKey))
|
||||||
|
if err != nil || strings.TrimSpace(raw) == "" {
|
||||||
|
return ids.UUID{}, err
|
||||||
|
}
|
||||||
|
conversationID, err := ids.Parse(strings.TrimSpace(raw))
|
||||||
|
if err != nil {
|
||||||
|
return ids.UUID{}, err
|
||||||
|
}
|
||||||
|
return conversationID, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (al *AgentLoop) persistConversationBinding(ctx context.Context, sessionKey string, conversationID ids.UUID) error {
|
||||||
|
if al == nil || al.memDelegate == nil || strings.TrimSpace(sessionKey) == "" || conversationID.IsZero() {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return al.memDelegate.UpsertKV(ctx, pkg.NAME, conversationBindingKey(sessionKey), conversationID.String())
|
||||||
|
}
|
||||||
|
|
||||||
type assembledContext struct {
|
type assembledContext struct {
|
||||||
systemPrompt string
|
systemPrompt string
|
||||||
userPrompt string
|
userPrompt string
|
||||||
|
|
@ -44,6 +73,20 @@ type agentRunMetrics struct {
|
||||||
TotalTokens int
|
TotalTokens int
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type ctxToolSessionKey struct{}
|
||||||
|
|
||||||
|
func withToolSessionKey(ctx context.Context, sessionKey string) context.Context {
|
||||||
|
if strings.TrimSpace(sessionKey) == "" {
|
||||||
|
return ctx
|
||||||
|
}
|
||||||
|
return context.WithValue(ctx, ctxToolSessionKey{}, sessionKey)
|
||||||
|
}
|
||||||
|
|
||||||
|
func toolSessionKeyFromContext(ctx context.Context) string {
|
||||||
|
v, _ := ctx.Value(ctxToolSessionKey{}).(string)
|
||||||
|
return strings.TrimSpace(v)
|
||||||
|
}
|
||||||
|
|
||||||
func collectAgentRunMetrics(result *fantasy.AgentResult) agentRunMetrics {
|
func collectAgentRunMetrics(result *fantasy.AgentResult) agentRunMetrics {
|
||||||
if result == nil {
|
if result == nil {
|
||||||
return agentRunMetrics{}
|
return agentRunMetrics{}
|
||||||
|
|
@ -82,13 +125,35 @@ func (al *AgentLoop) prepareRuntimeState(ctx context.Context, sessionKey string)
|
||||||
if cached, ok := al.conversationIDs.Load(sessionKey); ok {
|
if cached, ok := al.conversationIDs.Load(sessionKey); ok {
|
||||||
conversationID = cached
|
conversationID = cached
|
||||||
} else {
|
} else {
|
||||||
conversationID = ids.New()
|
boundID, err := al.loadBoundConversationID(ctx, sessionKey)
|
||||||
title := sessionKey
|
if err != nil {
|
||||||
if _, err := al.queries.CreateAgentConversation(ctx, memsqlc.CreateAgentConversationParams{
|
return ids.UUID{}, ids.UUID{}, fmt.Errorf("load conversation binding: %w", err)
|
||||||
ID: conversationID,
|
}
|
||||||
Title: &title,
|
if !boundID.IsZero() {
|
||||||
}); err != nil {
|
if _, err := al.queries.GetAgentConversation(ctx, memsqlc.GetAgentConversationParams{ID: boundID}); err == nil {
|
||||||
return ids.UUID{}, ids.UUID{}, fmt.Errorf("create agent conversation: %w", err)
|
conversationID = boundID
|
||||||
|
} else if err != sql.ErrNoRows {
|
||||||
|
return ids.UUID{}, ids.UUID{}, fmt.Errorf("load bound conversation: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if conversationID.IsZero() {
|
||||||
|
conversationID, err = al.lookupConversationIDForSession(ctx, sessionKey)
|
||||||
|
if err != nil {
|
||||||
|
if !errors.Is(err, sql.ErrNoRows) {
|
||||||
|
return ids.UUID{}, ids.UUID{}, fmt.Errorf("lookup conversation for session: %w", err)
|
||||||
|
}
|
||||||
|
conversationID = ids.New()
|
||||||
|
title := sessionKey
|
||||||
|
if _, err := al.queries.CreateAgentConversation(ctx, memsqlc.CreateAgentConversationParams{
|
||||||
|
ID: conversationID,
|
||||||
|
Title: &title,
|
||||||
|
}); err != nil {
|
||||||
|
return ids.UUID{}, ids.UUID{}, fmt.Errorf("create agent conversation: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := al.persistConversationBinding(ctx, sessionKey, conversationID); err != nil {
|
||||||
|
return ids.UUID{}, ids.UUID{}, fmt.Errorf("persist conversation binding: %w", err)
|
||||||
}
|
}
|
||||||
al.conversationIDs.Store(sessionKey, conversationID)
|
al.conversationIDs.Store(sessionKey, conversationID)
|
||||||
}
|
}
|
||||||
|
|
@ -150,15 +215,18 @@ func (al *AgentLoop) assembleContext(ctx context.Context, opts processOptions) (
|
||||||
builtMsgs := al.buildPromptMessages(opts, history, summary)
|
builtMsgs := al.buildPromptMessages(opts, history, summary)
|
||||||
systemPrompt, historyMsgs, userPrompt = al.splitMessages(opts, builtMsgs)
|
systemPrompt, historyMsgs, userPrompt = al.splitMessages(opts, builtMsgs)
|
||||||
}
|
}
|
||||||
if isPlanningOnlyPrompt(opts.UserMessage) {
|
if constraint := turnConstraintForQuery(opts.UserMessage); constraint != "" {
|
||||||
systemPrompt = strings.TrimSpace(systemPrompt + "\n\n## Turn Constraint\nThis request is planning-only. Answer directly in plain language. Do not call tools, do not emit tool-call syntax, and do not persist or schedule anything unless the user explicitly asked for that.")
|
systemPrompt = strings.TrimSpace(systemPrompt + "\n\n## Turn Constraint\n" + constraint)
|
||||||
} else if hintedNames := toolNames(al.initialPromptTools(opts.UserMessage)); len(hintedNames) > 0 {
|
}
|
||||||
systemPrompt = strings.TrimSpace(systemPrompt + fmt.Sprintf("\n\n## Turn Tool Hints\nFor this request, use these direct tools first: %s. This is an execution request, so do the tool work immediately instead of only describing intent. After the tool work finishes, always provide a concise final answer. If a tool fails or times out, explain that plainly in the final answer instead of stopping silently.", strings.Join(hintedNames, ", ")))
|
if !isPlanningOnlyPrompt(opts.UserMessage) {
|
||||||
if command := explicitExecCommand(opts.UserMessage); command != "" {
|
if hintedNames := initialPromptToolNames(al.tools, opts.UserMessage); len(hintedNames) > 0 {
|
||||||
systemPrompt = strings.TrimSpace(systemPrompt + fmt.Sprintf("\nIf you use exec for this request, set `command` to exactly %q. Do not substitute placeholders like `:`, empty strings, or paraphrases.", command))
|
systemPrompt = strings.TrimSpace(systemPrompt + "\n\n## Turn Tool Hints\n" + turnToolHintText(hintedNames))
|
||||||
}
|
if command := explicitExecCommand(opts.UserMessage); command != "" {
|
||||||
if skillName := explicitSkillName(opts.UserMessage); skillName != "" && strings.Contains(strings.ToLower(opts.UserMessage), "skill") {
|
systemPrompt = strings.TrimSpace(systemPrompt + fmt.Sprintf("\nIf you use exec for this request, set `command` to exactly %q. Do not substitute placeholders like `:`, empty strings, or paraphrases.", command))
|
||||||
systemPrompt = strings.TrimSpace(systemPrompt + fmt.Sprintf("\nIf you use skill_read for this request, set `name` to exactly %q. Do not substitute placeholders or punctuation-only values.", skillName))
|
}
|
||||||
|
if skillName := explicitSkillName(opts.UserMessage); skillName != "" && strings.Contains(strings.ToLower(opts.UserMessage), "skill") {
|
||||||
|
systemPrompt = strings.TrimSpace(systemPrompt + fmt.Sprintf("\nIf you use skill_read for this request, set `name` to exactly %q. Do not substitute placeholders or punctuation-only values.", skillName))
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -218,8 +286,6 @@ type ctxBlockCacheEntry struct {
|
||||||
const ctxBlockCacheTTL = 2 * time.Minute
|
const ctxBlockCacheTTL = 2 * time.Minute
|
||||||
|
|
||||||
func (al *AgentLoop) refreshContextBlocks(ctx context.Context, opts processOptions) {
|
func (al *AgentLoop) refreshContextBlocks(ctx context.Context, opts processOptions) {
|
||||||
al.updateToolContexts(opts.Channel, opts.ChatID)
|
|
||||||
|
|
||||||
var (
|
var (
|
||||||
obsBlock string
|
obsBlock string
|
||||||
kb string
|
kb string
|
||||||
|
|
@ -317,7 +383,7 @@ func (al *AgentLoop) loadSessionState(ctx context.Context, opts processOptions)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (al *AgentLoop) buildPromptMessages(opts processOptions, history []messages.Message, summary string) []messages.Message {
|
func (al *AgentLoop) buildPromptMessages(opts processOptions, history []messages.Message, summary string) []messages.Message {
|
||||||
builtMsgs := al.contextBuilder.BuildMessages(history, summary, opts.UserMessage, nil, opts.Channel, opts.ChatID)
|
builtMsgs := al.contextBuilder.BuildMessages(opts.SessionKey, history, summary, opts.UserMessage, nil, opts.Channel, opts.ChatID)
|
||||||
return builtMsgs
|
return builtMsgs
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -409,8 +475,24 @@ func (al *AgentLoop) initialPromptTools(query string) []tools.Tool {
|
||||||
if al == nil || al.tools == nil {
|
if al == nil || al.tools == nil {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
return collectInitialPromptTools(al.tools, query)
|
||||||
|
}
|
||||||
|
|
||||||
|
func initialPromptToolNames(registry *tools.ToolRegistry, query string) []string {
|
||||||
|
hinted := collectInitialPromptTools(registry, query)
|
||||||
|
return toolNames(hinted)
|
||||||
|
}
|
||||||
|
|
||||||
|
func collectInitialPromptTools(registry *tools.ToolRegistry, query string) []tools.Tool {
|
||||||
|
if registry == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
q := strings.ToLower(query)
|
q := strings.ToLower(query)
|
||||||
|
if tool := directDelegationTool(q); tool != "" {
|
||||||
|
return collectInitialTools(registry, query, map[string]struct{}{tool: {}})
|
||||||
|
}
|
||||||
|
|
||||||
want := map[string]struct{}{}
|
want := map[string]struct{}{}
|
||||||
|
|
||||||
if isToolDiscoveryPrompt(q) {
|
if isToolDiscoveryPrompt(q) {
|
||||||
|
|
@ -418,7 +500,7 @@ func (al *AgentLoop) initialPromptTools(query string) []tools.Tool {
|
||||||
if strings.Contains(q, "tool_call") {
|
if strings.Contains(q, "tool_call") {
|
||||||
want["tool_call"] = struct{}{}
|
want["tool_call"] = struct{}{}
|
||||||
}
|
}
|
||||||
return al.collectInitialTools(query, want)
|
return collectInitialTools(registry, query, want)
|
||||||
}
|
}
|
||||||
|
|
||||||
if strings.Contains(q, "skill") {
|
if strings.Contains(q, "skill") {
|
||||||
|
|
@ -453,11 +535,11 @@ func (al *AgentLoop) initialPromptTools(query string) []tools.Tool {
|
||||||
want["list_dir"] = struct{}{}
|
want["list_dir"] = struct{}{}
|
||||||
}
|
}
|
||||||
|
|
||||||
if strings.Contains(q, "spawn ") || strings.Contains(q, "background task") || strings.Contains(q, "async") {
|
if isExplicitSpawnPrompt(q) {
|
||||||
want["spawn"] = struct{}{}
|
want["spawn"] = struct{}{}
|
||||||
}
|
}
|
||||||
|
|
||||||
if strings.Contains(q, "subagent") || strings.Contains(q, "delegate") {
|
if isExplicitSubagentPrompt(q) {
|
||||||
want["subagent"] = struct{}{}
|
want["subagent"] = struct{}{}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -488,15 +570,22 @@ func (al *AgentLoop) initialPromptTools(query string) []tools.Tool {
|
||||||
want["web_fetch"] = struct{}{}
|
want["web_fetch"] = struct{}{}
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(want) == 0 {
|
if len(want) == 0 && shouldDefaultToToolSearch(q) {
|
||||||
want["tool_search"] = struct{}{}
|
want["tool_search"] = struct{}{}
|
||||||
}
|
}
|
||||||
|
|
||||||
return al.collectInitialTools(query, want)
|
return collectInitialTools(registry, query, want)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (al *AgentLoop) collectInitialTools(query string, want map[string]struct{}) []tools.Tool {
|
func (al *AgentLoop) collectInitialTools(query string, want map[string]struct{}) []tools.Tool {
|
||||||
if al == nil || al.tools == nil || len(want) == 0 {
|
if al == nil || al.tools == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return collectInitialTools(al.tools, query, want)
|
||||||
|
}
|
||||||
|
|
||||||
|
func collectInitialTools(registry *tools.ToolRegistry, query string, want map[string]struct{}) []tools.Tool {
|
||||||
|
if registry == nil || len(want) == 0 {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -525,7 +614,7 @@ func (al *AgentLoop) collectInitialTools(query string, want map[string]struct{})
|
||||||
if _, ok := want[name]; !ok {
|
if _, ok := want[name]; !ok {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
tool, found := al.tools.Get(name)
|
tool, found := registry.Get(name)
|
||||||
if !found {
|
if !found {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
@ -560,6 +649,158 @@ func shouldExposeToolResultSearch(q string) bool {
|
||||||
strings.Contains(query, "search tool results")
|
strings.Contains(query, "search tool results")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func directDelegationTool(query string) string {
|
||||||
|
q := strings.ToLower(strings.TrimSpace(query))
|
||||||
|
if q == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
if isExplicitSpawnPrompt(q) {
|
||||||
|
return "spawn"
|
||||||
|
}
|
||||||
|
if isExplicitSubagentPrompt(q) {
|
||||||
|
return "subagent"
|
||||||
|
}
|
||||||
|
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizeDelegationQuery(q string) string {
|
||||||
|
normalized := strings.ToLower(strings.TrimSpace(q))
|
||||||
|
for {
|
||||||
|
trimmed := normalized
|
||||||
|
for _, prefix := range []string{"please ", "can you ", "could you ", "would you ", "kindly "} {
|
||||||
|
if strings.HasPrefix(trimmed, prefix) {
|
||||||
|
trimmed = strings.TrimSpace(strings.TrimPrefix(trimmed, prefix))
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if trimmed == normalized {
|
||||||
|
return normalized
|
||||||
|
}
|
||||||
|
normalized = trimmed
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func hasImperativePrefix(q string, prefixes ...string) bool {
|
||||||
|
normalized := normalizeDelegationQuery(q)
|
||||||
|
if strings.Contains(normalized, "?") {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
for _, prefix := range prefixes {
|
||||||
|
if strings.HasPrefix(normalized, prefix) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
func isMetaExecutionDiscussionPrompt(q string) bool {
|
||||||
|
normalized := normalizeDelegationQuery(q)
|
||||||
|
if !(strings.Contains(normalized, "subagent") || strings.Contains(normalized, "delegate") || strings.Contains(normalized, "background") || strings.Contains(normalized, "asynchronously")) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if strings.Contains(normalized, "?") {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
if strings.Contains(normalized, " or ") {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
for _, prefix := range []string{"when should", "why ", "how ", "should we", "explain whether", "plan to ", "give me a plan", "whether ", "when ", "compare "} {
|
||||||
|
if strings.HasPrefix(normalized, prefix) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
func shouldDefaultToToolSearch(q string) bool {
|
||||||
|
normalized := normalizeDelegationQuery(q)
|
||||||
|
if normalized == "" || strings.Contains(normalized, "?") || isPlanningOnlyPrompt(normalized) || isMetaExecutionDiscussionPrompt(normalized) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
for _, prefix := range []string{"debug ", "review ", "inspect ", "investigate ", "analyze ", "analyse ", "fix ", "implement ", "trace ", "profile ", "audit ", "check ", "examine ", "look into ", "look at ", "compare ", "verify "} {
|
||||||
|
if strings.HasPrefix(normalized, prefix) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
func turnToolHintText(hintedNames []string) string {
|
||||||
|
if len(hintedNames) == 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
if len(hintedNames) == 1 && hintedNames[0] == "tool_search" {
|
||||||
|
return "For this request, if you need a tool, start with `tool_search` to discover the right concrete tool. Only promote a concrete tool when it clearly matches the user's request, and answer directly if no tool is needed."
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("For this request, use these direct tools first: %s. Treat these as the initially available direct tools for this turn; do not call unrelated tools unless they are explicitly promoted later. This is an execution request, so do the tool work immediately instead of only describing intent. After the tool work finishes, always provide a concise final answer. If a tool fails or times out, explain that plainly in the final answer instead of stopping silently.", strings.Join(hintedNames, ", "))
|
||||||
|
}
|
||||||
|
|
||||||
|
func isExplicitSpawnPrompt(q string) bool {
|
||||||
|
return hasImperativePrefix(q,
|
||||||
|
"spawn ",
|
||||||
|
"spawn a ",
|
||||||
|
"spawn an ",
|
||||||
|
"start a background task",
|
||||||
|
"start a background job",
|
||||||
|
"run this in the background",
|
||||||
|
"run it in the background",
|
||||||
|
"execute this in the background",
|
||||||
|
"execute it in the background",
|
||||||
|
"do this in the background",
|
||||||
|
"do it in the background",
|
||||||
|
"run this asynchronously",
|
||||||
|
"run it asynchronously",
|
||||||
|
"execute this asynchronously",
|
||||||
|
"execute it asynchronously",
|
||||||
|
"do this asynchronously",
|
||||||
|
"do it asynchronously",
|
||||||
|
"start an async task",
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
func isExplicitSubagentPrompt(q string) bool {
|
||||||
|
return hasImperativePrefix(q,
|
||||||
|
"use a subagent",
|
||||||
|
"use the subagent",
|
||||||
|
"use subagent",
|
||||||
|
"ask a subagent",
|
||||||
|
"have a subagent",
|
||||||
|
"delegate this",
|
||||||
|
"delegate it",
|
||||||
|
"delegate the task",
|
||||||
|
"delegate this task",
|
||||||
|
"delegate to a subagent",
|
||||||
|
"hand this off to a subagent",
|
||||||
|
"hand it off to a subagent",
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
func turnConstraintForQuery(query string) string {
|
||||||
|
if isPlanningOnlyPrompt(query) {
|
||||||
|
return "This request is planning-only. Answer directly in plain language. Do not call tools, do not emit tool-call syntax, and do not persist or schedule anything unless the user explicitly asked for that. Keep the answer compact and structured: no preamble, no recap, and no filler. Use the shortest format that fully covers the requested horizon, with brief day/week bullets and brief carry-forward or reminder notes only."
|
||||||
|
}
|
||||||
|
if tool := directDelegationTool(query); tool != "" {
|
||||||
|
return fmt.Sprintf("This request explicitly asks for delegated execution. Call `%s` as your first tool step. Do not use tool_search or tool_call first, and do not solve the task yourself before delegating. After the delegated result returns, answer with that result plainly and concisely.", tool)
|
||||||
|
}
|
||||||
|
if command := explicitExecCommand(query); command != "" {
|
||||||
|
return fmt.Sprintf("This request explicitly asks for command execution. Call `exec` as your first tool step with `command` set to exactly %q. After the tool returns, answer using the actual command output. Do not claim permission denial or failure unless the tool result says so.", command)
|
||||||
|
}
|
||||||
|
if path, content := explicitWriteFileRequest(query); path != "" {
|
||||||
|
constraint := fmt.Sprintf("This request explicitly asks for a file write. Call `write_file` as your first tool step with `path` set to exactly %q", path)
|
||||||
|
if content != "" {
|
||||||
|
constraint += fmt.Sprintf(" and `content` set to exactly %q", content)
|
||||||
|
}
|
||||||
|
constraint += ". Do not claim success unless the tool call succeeds."
|
||||||
|
return constraint
|
||||||
|
}
|
||||||
|
if isCommitmentCapturePrompt(query) {
|
||||||
|
return "This request explicitly asks you to capture commitments. Use `memory` to write each distinct commitment once, then stop calling tools and answer directly with a concise reminder/follow-up plan. Do not repeat the same memory write or loop on memory writes."
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
func isPlanningOnlyPrompt(query string) bool {
|
func isPlanningOnlyPrompt(query string) bool {
|
||||||
q := strings.ToLower(strings.TrimSpace(query))
|
q := strings.ToLower(strings.TrimSpace(query))
|
||||||
if q == "" {
|
if q == "" {
|
||||||
|
|
@ -607,6 +848,8 @@ func isPlanningOnlyPrompt(query string) bool {
|
||||||
"capture ",
|
"capture ",
|
||||||
"set reminder",
|
"set reminder",
|
||||||
"remind me",
|
"remind me",
|
||||||
|
"schedule reminder",
|
||||||
|
"create reminder",
|
||||||
"subagent",
|
"subagent",
|
||||||
"spawn ",
|
"spawn ",
|
||||||
"background task",
|
"background task",
|
||||||
|
|
@ -620,6 +863,21 @@ func isPlanningOnlyPrompt(query string) bool {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func isCommitmentCapturePrompt(query string) bool {
|
||||||
|
q := strings.ToLower(strings.TrimSpace(query))
|
||||||
|
if q == "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !(strings.Contains(q, "commitment") || strings.Contains(q, "commitments")) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return strings.Contains(q, "capture") ||
|
||||||
|
strings.Contains(q, "track") ||
|
||||||
|
strings.Contains(q, "remember") ||
|
||||||
|
strings.Contains(q, "store") ||
|
||||||
|
strings.Contains(q, "i have")
|
||||||
|
}
|
||||||
|
|
||||||
func (al *AgentLoop) createFantasyAgent(ctx context.Context, opts processOptions, systemPrompt string, adaptedTools []fantasy.AgentTool, prepareStep func(context.Context, fantasy.PrepareStepFunctionOptions) (context.Context, fantasy.PrepareStepResult, error)) (fantasy.Agent, ids.UUID, ids.UUID, error) {
|
func (al *AgentLoop) createFantasyAgent(ctx context.Context, opts processOptions, systemPrompt string, adaptedTools []fantasy.AgentTool, prepareStep func(context.Context, fantasy.PrepareStepFunctionOptions) (context.Context, fantasy.PrepareStepResult, error)) (fantasy.Agent, ids.UUID, ids.UUID, error) {
|
||||||
conversationID, runID, err := al.prepareRuntimeState(ctx, opts.SessionKey)
|
conversationID, runID, err := al.prepareRuntimeState(ctx, opts.SessionKey)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|
@ -635,12 +893,15 @@ func (al *AgentLoop) createFantasyAgent(ctx context.Context, opts processOptions
|
||||||
ThresholdChars: al.offloadThresholdChars,
|
ThresholdChars: al.offloadThresholdChars,
|
||||||
}
|
}
|
||||||
toolRuntime := SecureBusToolRuntime{
|
toolRuntime := SecureBusToolRuntime{
|
||||||
Base: baseRuntime,
|
Offloader: baseRuntime,
|
||||||
Bus: al.secureBus,
|
FantasyTools: fantasyToolMap(adaptedTools, al.tools),
|
||||||
SessionKey: opts.SessionKey,
|
Bus: al.secureBus,
|
||||||
UserPrompt: opts.UserMessage,
|
SessionKey: opts.SessionKey,
|
||||||
StateStore: al.stateStore,
|
Channel: opts.Channel,
|
||||||
RunID: runID,
|
ChatID: opts.ChatID,
|
||||||
|
UserPrompt: opts.UserMessage,
|
||||||
|
StateStore: al.stateStore,
|
||||||
|
RunID: runID,
|
||||||
}
|
}
|
||||||
transitionObserver := fantasy.ReActTransitionObserverFunc(func(observerCtx context.Context, t fantasy.ReActTransition) {
|
transitionObserver := fantasy.ReActTransitionObserverFunc(func(observerCtx context.Context, t fantasy.ReActTransition) {
|
||||||
if al.stateStore == nil || runID.IsZero() {
|
if al.stateStore == nil || runID.IsZero() {
|
||||||
|
|
@ -811,33 +1072,74 @@ func (al *AgentLoop) resolveFinalContent(finalContent string, steps []fantasy.St
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if best.text != "" && best.score > 0 {
|
|
||||||
logger.WarnCF("agent", "Recovered empty final response",
|
|
||||||
map[string]interface{}{
|
|
||||||
"candidates": len(candidates),
|
|
||||||
"score": best.score,
|
|
||||||
"source": best.source,
|
|
||||||
})
|
|
||||||
return best.text, nil
|
|
||||||
}
|
|
||||||
if best.text != "" {
|
|
||||||
logger.WarnCF("agent", "Recovered empty final response from fallback tool content",
|
|
||||||
map[string]interface{}{
|
|
||||||
"candidates": len(candidates),
|
|
||||||
"score": best.score,
|
|
||||||
"source": best.source,
|
|
||||||
})
|
|
||||||
return best.text, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
toolCalls := 0
|
toolCalls := 0
|
||||||
for _, step := range steps {
|
for _, step := range steps {
|
||||||
toolCalls += len(step.Content.ToolCalls())
|
toolCalls += len(step.Content.ToolCalls())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if best.text != "" && best.score > 0 {
|
||||||
|
if toolCalls == 0 && best.source == "step_text" && best.score <= 1 {
|
||||||
|
if fallback := fallbackNoToolResponse(al, steps); fallback != "" {
|
||||||
|
return fallback, nil
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
logger.WarnCF("agent", "Recovered empty final response",
|
||||||
|
map[string]interface{}{
|
||||||
|
"candidates": len(candidates),
|
||||||
|
"score": best.score,
|
||||||
|
"source": best.source,
|
||||||
|
})
|
||||||
|
return best.text, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if best.text != "" {
|
||||||
|
if toolCalls == 0 {
|
||||||
|
if fallback := fallbackNoToolResponse(al, steps); fallback != "" {
|
||||||
|
return fallback, nil
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
logger.WarnCF("agent", "Recovered empty final response from fallback tool content",
|
||||||
|
map[string]interface{}{
|
||||||
|
"candidates": len(candidates),
|
||||||
|
"score": best.score,
|
||||||
|
"source": best.source,
|
||||||
|
})
|
||||||
|
return best.text, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if toolCalls == 0 {
|
||||||
|
if fallback := fallbackNoToolResponse(al, steps); fallback != "" {
|
||||||
|
return fallback, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
return "", fmt.Errorf("agent produced no final response text (steps=%d, tool_calls=%d)", len(steps), toolCalls)
|
return "", fmt.Errorf("agent produced no final response text (steps=%d, tool_calls=%d)", len(steps), toolCalls)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func fallbackNoToolResponse(al *AgentLoop, steps []fantasy.StepResult) string {
|
||||||
|
if len(steps) != 1 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
stepText := strings.TrimSpace(steps[0].Content.Text())
|
||||||
|
if stepText == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
lower := strings.ToLower(stepText)
|
||||||
|
if strings.Contains(lower, "which file") || strings.Contains(lower, "what do you want") || strings.Contains(lower, "what would you like") {
|
||||||
|
return stepText
|
||||||
|
}
|
||||||
|
if al != nil {
|
||||||
|
if grounded := strings.TrimSpace(al.groundFinalContent(stepText, stepText, steps)); grounded != "" {
|
||||||
|
if grounded != stepText {
|
||||||
|
return grounded
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
func (al *AgentLoop) groundFinalContent(userPrompt, finalContent string, steps []fantasy.StepResult) string {
|
func (al *AgentLoop) groundFinalContent(userPrompt, finalContent string, steps []fantasy.StepResult) string {
|
||||||
grounded := strings.TrimSpace(finalContent)
|
grounded := strings.TrimSpace(finalContent)
|
||||||
if grounded == "" {
|
if grounded == "" {
|
||||||
|
|
@ -896,6 +1198,11 @@ func (al *AgentLoop) groundFinalContent(userPrompt, finalContent string, steps [
|
||||||
(!mentionsExecFailure(lowerFinal) || strings.Contains(lowerFinal, "completed successfully")) {
|
(!mentionsExecFailure(lowerFinal) || strings.Contains(lowerFinal, "completed successfully")) {
|
||||||
return execError
|
return execError
|
||||||
}
|
}
|
||||||
|
if execOutput := detectExecSuccessText(toolTexts); execOutput != "" &&
|
||||||
|
asksForExecResult(lowerPrompt) &&
|
||||||
|
mentionsExecFailure(lowerFinal) {
|
||||||
|
return formatExecSuccessResponse(execOutput)
|
||||||
|
}
|
||||||
|
|
||||||
if strings.Contains(lowerPrompt, "commitment") || strings.Contains(lowerPrompt, "commitments") {
|
if strings.Contains(lowerPrompt, "commitment") || strings.Contains(lowerPrompt, "commitments") {
|
||||||
clauses := extractCommitmentClauses(userPrompt)
|
clauses := extractCommitmentClauses(userPrompt)
|
||||||
|
|
@ -930,6 +1237,24 @@ func (al *AgentLoop) groundFinalContent(userPrompt, finalContent string, steps [
|
||||||
!strings.Contains(lowerFinal, "timeline") {
|
!strings.Contains(lowerFinal, "timeline") {
|
||||||
return strings.TrimSpace(grounded + "\n\nReminder/follow-up plan: schedule each item against its stated timing and review progress at each checkpoint.")
|
return strings.TrimSpace(grounded + "\n\nReminder/follow-up plan: schedule each item against its stated timing and review progress at each checkpoint.")
|
||||||
}
|
}
|
||||||
|
if strings.Contains(lowerPrompt, "daily plan") &&
|
||||||
|
strings.Contains(lowerPrompt, "carry forward") &&
|
||||||
|
!strings.Contains(lowerFinal, "monday") &&
|
||||||
|
!strings.Contains(lowerFinal, "tuesday") {
|
||||||
|
return strings.TrimSpace(expandWeekdayAbbreviations(grounded))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if strings.Contains(lowerPrompt, "webinar") &&
|
||||||
|
strings.Contains(lowerPrompt, "risk") &&
|
||||||
|
strings.Contains(lowerPrompt, "follow-up") {
|
||||||
|
if (strings.Contains(lowerFinal, "risk") || strings.Contains(lowerFinal, "fallback") || strings.Contains(lowerFinal, "failure")) &&
|
||||||
|
!strings.Contains(lowerFinal, "follow-up") &&
|
||||||
|
!strings.Contains(lowerFinal, "follow up") &&
|
||||||
|
!strings.Contains(lowerFinal, "check-in") &&
|
||||||
|
!strings.Contains(lowerFinal, "verify") {
|
||||||
|
return strings.TrimSpace(grounded + "\n\nFollow-up actions: schedule a 24-hour follow-up check-in to send the recording and slides, verify attendee follow-up status, and review the risk/fallback notes before the next webinar.")
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if strings.Contains(lowerPrompt, "skill") &&
|
if strings.Contains(lowerPrompt, "skill") &&
|
||||||
|
|
@ -961,6 +1286,26 @@ func (al *AgentLoop) groundFinalContent(userPrompt, finalContent string, steps [
|
||||||
return grounded
|
return grounded
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func expandWeekdayAbbreviations(text string) string {
|
||||||
|
replacer := strings.NewReplacer(
|
||||||
|
"**Mon", "**Monday",
|
||||||
|
"**Tue", "**Tuesday",
|
||||||
|
"**Wed", "**Wednesday",
|
||||||
|
"**Thu", "**Thursday",
|
||||||
|
"**Fri", "**Friday",
|
||||||
|
"**Sat", "**Saturday",
|
||||||
|
"**Sun", "**Sunday",
|
||||||
|
" Mon ", " Monday ",
|
||||||
|
" Tue ", " Tuesday ",
|
||||||
|
" Wed ", " Wednesday ",
|
||||||
|
" Thu ", " Thursday ",
|
||||||
|
" Fri ", " Friday ",
|
||||||
|
" Sat ", " Saturday ",
|
||||||
|
" Sun ", " Sunday ",
|
||||||
|
)
|
||||||
|
return replacer.Replace(text)
|
||||||
|
}
|
||||||
|
|
||||||
func collectToolTexts(steps []fantasy.StepResult) map[string][]string {
|
func collectToolTexts(steps []fantasy.StepResult) map[string][]string {
|
||||||
toolTexts := make(map[string][]string)
|
toolTexts := make(map[string][]string)
|
||||||
for _, step := range steps {
|
for _, step := range steps {
|
||||||
|
|
@ -1089,9 +1434,15 @@ func explicitWriteFileRequest(prompt string) (string, string) {
|
||||||
}
|
}
|
||||||
|
|
||||||
path := ""
|
path := ""
|
||||||
pathRE := regexp.MustCompile(`(?i)(?:to a file called|file called)\s+([^\s"'` + "`" + `,]+)`)
|
pathRE := regexp.MustCompile("(?i)(?:to a file called|file called)\\s+(?:\"([^\"]+)\"|'([^']+)'|`([^`]+)`|([^\\s\"'`,]+))")
|
||||||
if matches := pathRE.FindStringSubmatch(prompt); len(matches) > 1 {
|
if matches := pathRE.FindStringSubmatch(prompt); len(matches) > 1 {
|
||||||
path = strings.Trim(matches[1], "\"'`.,")
|
for _, candidate := range matches[1:] {
|
||||||
|
candidate = strings.TrimSpace(candidate)
|
||||||
|
if candidate != "" {
|
||||||
|
path = strings.Trim(candidate, "\"'`.,")
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
content := ""
|
content := ""
|
||||||
|
|
@ -1220,6 +1571,31 @@ func detectExecErrorText(toolTexts map[string][]string) string {
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func detectExecSuccessText(toolTexts map[string][]string) string {
|
||||||
|
if detectExecErrorText(toolTexts) != "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
for i := len(toolTexts["exec"]) - 1; i >= 0; i-- {
|
||||||
|
text := strings.TrimSpace(toolTexts["exec"][i])
|
||||||
|
if text == "" || text == "(no output)" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
return text
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
func formatExecSuccessResponse(output string) string {
|
||||||
|
trimmed := strings.TrimSpace(output)
|
||||||
|
if trimmed == "" {
|
||||||
|
return trimmed
|
||||||
|
}
|
||||||
|
if !strings.Contains(trimmed, "\n") {
|
||||||
|
return fmt.Sprintf("The output is:\n\n```\n%s\n```", trimmed)
|
||||||
|
}
|
||||||
|
return trimmed
|
||||||
|
}
|
||||||
|
|
||||||
func asksForMemorySearch(lowerPrompt string) bool {
|
func asksForMemorySearch(lowerPrompt string) bool {
|
||||||
return strings.Contains(lowerPrompt, "search your memory") ||
|
return strings.Contains(lowerPrompt, "search your memory") ||
|
||||||
(strings.Contains(lowerPrompt, "memory") && strings.Contains(lowerPrompt, "search")) ||
|
(strings.Contains(lowerPrompt, "memory") && strings.Contains(lowerPrompt, "search")) ||
|
||||||
|
|
@ -1248,6 +1624,8 @@ func mentionsNoResults(lowerText string) bool {
|
||||||
func mentionsExecFailure(lowerText string) bool {
|
func mentionsExecFailure(lowerText string) bool {
|
||||||
return strings.Contains(lowerText, "timed out") ||
|
return strings.Contains(lowerText, "timed out") ||
|
||||||
strings.Contains(lowerText, "blocked") ||
|
strings.Contains(lowerText, "blocked") ||
|
||||||
|
strings.Contains(lowerText, "denied") ||
|
||||||
|
strings.Contains(lowerText, "permission") ||
|
||||||
strings.Contains(lowerText, "cannot") ||
|
strings.Contains(lowerText, "cannot") ||
|
||||||
strings.Contains(lowerText, "placeholder") ||
|
strings.Contains(lowerText, "placeholder") ||
|
||||||
strings.Contains(lowerText, "not allowed") ||
|
strings.Contains(lowerText, "not allowed") ||
|
||||||
|
|
@ -1377,13 +1755,37 @@ func (al *AgentLoop) recoverSkillSummary(skillName string) string {
|
||||||
// runAgentLoop is the core message processing logic.
|
// runAgentLoop is the core message processing logic.
|
||||||
// It delegates to assembleContext for shared pre-processing, then branches on
|
// It delegates to assembleContext for shared pre-processing, then branches on
|
||||||
// opts.Streaming to either Generate (synchronous) or Stream (real-time deltas).
|
// opts.Streaming to either Generate (synchronous) or Stream (real-time deltas).
|
||||||
func (al *AgentLoop) runAgentLoop(ctx context.Context, opts processOptions) (string, error) {
|
func (al *AgentLoop) runAgentLoop(ctx context.Context, opts processOptions) (finalContent string, err error) {
|
||||||
|
ctx = withToolSessionKey(ctx, opts.SessionKey)
|
||||||
al.activeSessionKey.Store(opts.SessionKey)
|
al.activeSessionKey.Store(opts.SessionKey)
|
||||||
|
defer al.activeSessionKey.Store("")
|
||||||
|
defer func() {
|
||||||
|
if err == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if !opts.RunID.IsZero() {
|
||||||
|
al.persistFailedRun(ctx, opts, err)
|
||||||
|
if endErr := al.endTask(ctx, opts.ConversationID, opts.RunID, TaskCompletion{
|
||||||
|
TaskID: opts.SessionKey,
|
||||||
|
Description: utils.Truncate(opts.UserMessage, 100),
|
||||||
|
Completed: false,
|
||||||
|
CreatedAt: time.Now().UTC(),
|
||||||
|
}); endErr != nil {
|
||||||
|
logger.WarnCF("agent", "Failed to record failed task completion",
|
||||||
|
map[string]interface{}{"error": endErr.Error(), "session": opts.SessionKey})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if opts.SessionKey != "" {
|
||||||
|
go al.sessions.Save(opts.SessionKey)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
ac, err := al.assembleContext(ctx, opts)
|
ac, err := al.assembleContext(ctx, opts)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", err
|
return "", err
|
||||||
}
|
}
|
||||||
|
opts.ConversationID = ac.conversationID
|
||||||
|
opts.RunID = ac.runID
|
||||||
|
|
||||||
if opts.Streaming {
|
if opts.Streaming {
|
||||||
return al.runStreaming(ctx, opts, ac)
|
return al.runStreaming(ctx, opts, ac)
|
||||||
|
|
@ -1407,7 +1809,7 @@ func (al *AgentLoop) runAgentLoop(ctx context.Context, opts processOptions) (str
|
||||||
al.auditStep(ctx, step, opts.SessionKey)
|
al.auditStep(ctx, step, opts.SessionKey)
|
||||||
}
|
}
|
||||||
|
|
||||||
finalContent, err := al.resolveFinalContent(result.Response.Content.Text(), result.Steps)
|
finalContent, err = al.resolveFinalContent(result.Response.Content.Text(), result.Steps)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.ErrorCF("agent", "Agent finished without final response text",
|
logger.ErrorCF("agent", "Agent finished without final response text",
|
||||||
map[string]interface{}{
|
map[string]interface{}{
|
||||||
|
|
@ -1418,21 +1820,24 @@ func (al *AgentLoop) runAgentLoop(ctx context.Context, opts processOptions) (str
|
||||||
}
|
}
|
||||||
finalContent = al.groundFinalContent(opts.UserMessage, finalContent, result.Steps)
|
finalContent = al.groundFinalContent(opts.UserMessage, finalContent, result.Steps)
|
||||||
|
|
||||||
// Populate IDs for task completion tracking
|
|
||||||
opts.ConversationID = ac.conversationID
|
|
||||||
opts.RunID = ac.runID
|
|
||||||
|
|
||||||
return al.postProcess(ctx, opts, finalContent, collectAgentRunMetrics(result)), nil
|
return al.postProcess(ctx, opts, finalContent, collectAgentRunMetrics(result)), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// runStreaming uses Fantasy's agent.Stream() to stream token deltas to the bus
|
// runStreaming uses Fantasy's agent.Stream() to stream token deltas to the bus
|
||||||
// in real time, using the pre-assembled context from assembleContext.
|
// in real time, using the pre-assembled context from assembleContext.
|
||||||
func (al *AgentLoop) runStreaming(ctx context.Context, opts processOptions, ac assembledContext) (string, error) {
|
// Failure terminalization is owned by runAgentLoop so streaming errors only
|
||||||
|
// record terminal state once.
|
||||||
|
func (al *AgentLoop) runStreaming(ctx context.Context, opts processOptions, ac assembledContext) (finalContent string, err error) {
|
||||||
|
opts.ConversationID = ac.conversationID
|
||||||
|
opts.RunID = ac.runID
|
||||||
|
var streamedText strings.Builder
|
||||||
|
|
||||||
streamCall := fantasy.AgentStreamCall{
|
streamCall := fantasy.AgentStreamCall{
|
||||||
Prompt: ac.userPrompt,
|
Prompt: ac.userPrompt,
|
||||||
Messages: ac.fantasyHistory,
|
Messages: ac.fantasyHistory,
|
||||||
|
|
||||||
OnTextDelta: func(id, text string) error {
|
OnTextDelta: func(id, text string) error {
|
||||||
|
streamedText.WriteString(text)
|
||||||
if opts.Channel != "" && opts.ChatID != "" {
|
if opts.Channel != "" && opts.ChatID != "" {
|
||||||
al.bus.PublishOutbound(bus.OutboundMessage{
|
al.bus.PublishOutbound(bus.OutboundMessage{
|
||||||
Channel: opts.Channel,
|
Channel: opts.Channel,
|
||||||
|
|
@ -1470,7 +1875,11 @@ func (al *AgentLoop) runStreaming(ctx context.Context, opts processOptions, ac a
|
||||||
return "", fmt.Errorf("agent Stream failed: %w", err)
|
return "", fmt.Errorf("agent Stream failed: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
finalContent, err := al.resolveFinalContent(result.Response.Content.Text(), result.Steps)
|
responseText := result.Response.Content.Text()
|
||||||
|
if strings.TrimSpace(responseText) == "" {
|
||||||
|
responseText = streamedText.String()
|
||||||
|
}
|
||||||
|
finalContent, err = al.resolveFinalContent(responseText, result.Steps)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.ErrorCF("agent", "Streaming agent finished without final response text",
|
logger.ErrorCF("agent", "Streaming agent finished without final response text",
|
||||||
map[string]interface{}{
|
map[string]interface{}{
|
||||||
|
|
@ -1481,10 +1890,6 @@ func (al *AgentLoop) runStreaming(ctx context.Context, opts processOptions, ac a
|
||||||
}
|
}
|
||||||
finalContent = al.groundFinalContent(opts.UserMessage, finalContent, result.Steps)
|
finalContent = al.groundFinalContent(opts.UserMessage, finalContent, result.Steps)
|
||||||
|
|
||||||
// Populate IDs for task completion tracking
|
|
||||||
opts.ConversationID = ac.conversationID
|
|
||||||
opts.RunID = ac.runID
|
|
||||||
|
|
||||||
return al.postProcess(ctx, opts, finalContent, collectAgentRunMetrics(result)), nil
|
return al.postProcess(ctx, opts, finalContent, collectAgentRunMetrics(result)), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -1591,23 +1996,3 @@ func (al *AgentLoop) enqueueAuditEntry(entry *memory.AuditEntry) (ok bool) {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// updateToolContexts updates the context for tools that need channel/chatID info.
|
|
||||||
func (al *AgentLoop) updateToolContexts(channel, chatID string) {
|
|
||||||
// Use ContextualTool interface instead of type assertions
|
|
||||||
if tool, ok := al.tools.Get("message"); ok {
|
|
||||||
if mt, ok := tool.(tools.ContextualTool); ok {
|
|
||||||
mt.SetContext(channel, chatID)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if tool, ok := al.tools.Get("spawn"); ok {
|
|
||||||
if st, ok := tool.(tools.ContextualTool); ok {
|
|
||||||
st.SetContext(channel, chatID)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if tool, ok := al.tools.Get("subagent"); ok {
|
|
||||||
if st, ok := tool.(tools.ContextualTool); ok {
|
|
||||||
st.SetContext(channel, chatID)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
|
||||||
|
|
@ -2,6 +2,7 @@ package agent
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
@ -34,6 +35,7 @@ func (al *AgentLoop) persistRunCheckpoint(ctx context.Context, opts processOptio
|
||||||
if opts.ConversationID.IsZero() || opts.RunID.IsZero() {
|
if opts.ConversationID.IsZero() || opts.RunID.IsZero() {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
persistCtx := context.WithoutCancel(ctx)
|
||||||
|
|
||||||
history := al.sessions.GetHistory(opts.SessionKey)
|
history := al.sessions.GetHistory(opts.SessionKey)
|
||||||
snapshot := conversations.NewCheckpointSnapshot(
|
snapshot := conversations.NewCheckpointSnapshot(
|
||||||
|
|
@ -55,7 +57,7 @@ func (al *AgentLoop) persistRunCheckpoint(ctx context.Context, opts processOptio
|
||||||
"errors": metrics.Errors,
|
"errors": metrics.Errors,
|
||||||
}
|
}
|
||||||
|
|
||||||
runState, err := al.stateStore.AddRunState(ctx, opts.RunID, metrics.StepCount, fantasy.ReActStateDone, snapshot)
|
runState, err := al.stateStore.AddRunState(persistCtx, opts.RunID, metrics.StepCount, fantasy.ReActStateDone, snapshot)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.WarnCF("agent", "Failed to persist checkpointable run snapshot", map[string]any{
|
logger.WarnCF("agent", "Failed to persist checkpointable run snapshot", map[string]any{
|
||||||
"session_key": opts.SessionKey,
|
"session_key": opts.SessionKey,
|
||||||
|
|
@ -68,7 +70,7 @@ func (al *AgentLoop) persistRunCheckpoint(ctx context.Context, opts processOptio
|
||||||
meta["run_state_id"] = runState.ID.String()
|
meta["run_state_id"] = runState.ID.String()
|
||||||
|
|
||||||
checkpointStore := NewCheckpointStore(al.queries)
|
checkpointStore := NewCheckpointStore(al.queries)
|
||||||
if _, err := checkpointStore.CreateCheckpoint(ctx, opts.ConversationID, checkpointName, runState.ID, meta); err != nil {
|
if _, err := checkpointStore.CreateCheckpoint(persistCtx, opts.ConversationID, checkpointName, runState.ID, meta); err != nil {
|
||||||
logger.WarnCF("agent", "Failed to create runtime checkpoint", map[string]any{
|
logger.WarnCF("agent", "Failed to create runtime checkpoint", map[string]any{
|
||||||
"session_key": opts.SessionKey,
|
"session_key": opts.SessionKey,
|
||||||
"run_id": opts.RunID.String(),
|
"run_id": opts.RunID.String(),
|
||||||
|
|
@ -78,7 +80,7 @@ func (al *AgentLoop) persistRunCheckpoint(ctx context.Context, opts processOptio
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if _, err := al.stateStore.UpdateRunStatus(ctx, opts.RunID, "completed", meta); err != nil {
|
if _, err := al.stateStore.UpdateRunStatus(persistCtx, opts.RunID, "completed", meta); err != nil {
|
||||||
logger.WarnCF("agent", "Failed to update run completion status", map[string]any{
|
logger.WarnCF("agent", "Failed to update run completion status", map[string]any{
|
||||||
"session_key": opts.SessionKey,
|
"session_key": opts.SessionKey,
|
||||||
"run_id": opts.RunID.String(),
|
"run_id": opts.RunID.String(),
|
||||||
|
|
@ -87,6 +89,43 @@ func (al *AgentLoop) persistRunCheckpoint(ctx context.Context, opts processOptio
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (al *AgentLoop) persistFailedRun(ctx context.Context, opts processOptions, reason error) {
|
||||||
|
if al == nil || al.stateStore == nil || opts.RunID.IsZero() {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
persistCtx := context.WithoutCancel(ctx)
|
||||||
|
meta := map[string]any{
|
||||||
|
"session_key": opts.SessionKey,
|
||||||
|
}
|
||||||
|
if !opts.ConversationID.IsZero() {
|
||||||
|
meta["conversation_id"] = opts.ConversationID.String()
|
||||||
|
}
|
||||||
|
if reason != nil {
|
||||||
|
meta["error"] = reason.Error()
|
||||||
|
meta["reason"] = classifyRunFailure(reason)
|
||||||
|
}
|
||||||
|
if _, err := al.stateStore.UpdateRunStatus(persistCtx, opts.RunID, "failed", meta); err != nil {
|
||||||
|
logger.WarnCF("agent", "Failed to update run failure status", map[string]any{
|
||||||
|
"session_key": opts.SessionKey,
|
||||||
|
"run_id": opts.RunID.String(),
|
||||||
|
"error": err.Error(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func classifyRunFailure(err error) string {
|
||||||
|
if err == nil {
|
||||||
|
return "failed"
|
||||||
|
}
|
||||||
|
if errors.Is(err, context.Canceled) {
|
||||||
|
return "canceled"
|
||||||
|
}
|
||||||
|
if errors.Is(err, context.DeadlineExceeded) {
|
||||||
|
return "deadline_exceeded"
|
||||||
|
}
|
||||||
|
return "failed"
|
||||||
|
}
|
||||||
|
|
||||||
func (al *AgentLoop) RestoreSessionFromCheckpoint(ctx context.Context, sessionKey, checkpointName string) error {
|
func (al *AgentLoop) RestoreSessionFromCheckpoint(ctx context.Context, sessionKey, checkpointName string) error {
|
||||||
sessionKey = strings.TrimSpace(sessionKey)
|
sessionKey = strings.TrimSpace(sessionKey)
|
||||||
if sessionKey == "" {
|
if sessionKey == "" {
|
||||||
|
|
@ -112,6 +151,9 @@ func (al *AgentLoop) RestoreSessionFromCheckpoint(ctx context.Context, sessionKe
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
al.conversationIDs.Store(sessionKey, conversationID)
|
al.conversationIDs.Store(sessionKey, conversationID)
|
||||||
|
if err := al.persistConversationBinding(ctx, sessionKey, conversationID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -154,6 +196,9 @@ func (al *AgentLoop) ForkSessionFromCheckpoint(ctx context.Context, sourceSessio
|
||||||
return ids.UUID{}, err
|
return ids.UUID{}, err
|
||||||
}
|
}
|
||||||
al.conversationIDs.Store(forkSessionKey, conv.ID)
|
al.conversationIDs.Store(forkSessionKey, conv.ID)
|
||||||
|
if err := al.persistConversationBinding(ctx, forkSessionKey, conv.ID); err != nil {
|
||||||
|
return ids.UUID{}, err
|
||||||
|
}
|
||||||
return conv.ID, nil
|
return conv.ID, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -168,12 +213,23 @@ func (al *AgentLoop) lookupConversationIDForSession(ctx context.Context, session
|
||||||
if cached, ok := al.conversationIDs.Load(sessionKey); ok {
|
if cached, ok := al.conversationIDs.Load(sessionKey); ok {
|
||||||
return cached, nil
|
return cached, nil
|
||||||
}
|
}
|
||||||
|
if boundID, err := al.loadBoundConversationID(ctx, sessionKey); err == nil && !boundID.IsZero() {
|
||||||
|
if _, err := al.queries.GetAgentConversation(ctx, memsqlc.GetAgentConversationParams{ID: boundID}); err == nil {
|
||||||
|
al.conversationIDs.Store(sessionKey, boundID)
|
||||||
|
return boundID, nil
|
||||||
|
}
|
||||||
|
} else if err != nil {
|
||||||
|
return ids.UUID{}, fmt.Errorf("lookup conversation binding for session %q: %w", sessionKey, err)
|
||||||
|
}
|
||||||
|
|
||||||
conv, err := al.queries.GetLatestAgentConversationByTitle(ctx, memsqlc.GetLatestAgentConversationByTitleParams{Title: &sessionKey})
|
conv, err := al.queries.GetLatestAgentConversationByTitle(ctx, memsqlc.GetLatestAgentConversationByTitleParams{Title: &sessionKey})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return ids.UUID{}, fmt.Errorf("lookup conversation for session %q: %w", sessionKey, err)
|
return ids.UUID{}, fmt.Errorf("lookup conversation for session %q: %w", sessionKey, err)
|
||||||
}
|
}
|
||||||
al.conversationIDs.Store(sessionKey, conv.ID)
|
al.conversationIDs.Store(sessionKey, conv.ID)
|
||||||
|
if err := al.persistConversationBinding(ctx, sessionKey, conv.ID); err != nil {
|
||||||
|
return ids.UUID{}, fmt.Errorf("persist conversation binding for session %q: %w", sessionKey, err)
|
||||||
|
}
|
||||||
return conv.ID, nil
|
return conv.ID, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -121,6 +121,139 @@ func TestAgentLoop_ForkSessionFromCheckpoint_CreatesHydratedChildSession(t *test
|
||||||
assert.Equal(t, checkpointHistoryView(expected), agentMessageHistoryView(seeded))
|
assert.Equal(t, checkpointHistoryView(expected), agentMessageHistoryView(seeded))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestAgentLoop_RestartReusesConversationBinding(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
tmpDir, err := os.MkdirTemp("", "agent-checkpoint-restart-*")
|
||||||
|
require.NoError(t, err)
|
||||||
|
t.Cleanup(func() { _ = os.RemoveAll(tmpDir) })
|
||||||
|
|
||||||
|
cfg := config.DefaultConfig()
|
||||||
|
cfg.Agents.Defaults.Sandbox = tmpDir
|
||||||
|
cfg.Agents.Defaults.MaxTokens = 4096
|
||||||
|
cfg.Agents.Defaults.MaxToolIterations = 4
|
||||||
|
cfg.Memory.DBPath = filepath.Join(tmpDir, "agent-restart.db")
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(t.Context(), 10*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
al1 := mustNewAgentLoop(t, cfg, bus.NewMessageBus(), newMockLanguageModel("first response"))
|
||||||
|
sessionKey := "restart-session"
|
||||||
|
_, err = al1.processMessage(ctx, bus.InboundMessage{
|
||||||
|
Channel: "test",
|
||||||
|
SenderID: "user1",
|
||||||
|
ChatID: "chat1",
|
||||||
|
Content: "first turn",
|
||||||
|
SessionKey: sessionKey,
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
firstConversationID, err := al1.lookupConversationIDForSession(ctx, sessionKey)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, al1.sessions.Save(sessionKey))
|
||||||
|
al1.sessions.Close()
|
||||||
|
|
||||||
|
al2 := mustNewAgentLoop(t, cfg, bus.NewMessageBus(), newMockLanguageModel("second response"))
|
||||||
|
_, err = al2.processMessage(ctx, bus.InboundMessage{
|
||||||
|
Channel: "test",
|
||||||
|
SenderID: "user1",
|
||||||
|
ChatID: "chat1",
|
||||||
|
Content: "second turn",
|
||||||
|
SessionKey: sessionKey,
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
secondConversationID, err := al2.lookupConversationIDForSession(ctx, sessionKey)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
assert.Equal(t, firstConversationID, secondConversationID)
|
||||||
|
conversations, err := al2.queries.ListAgentConversations(ctx, memsqlc.ListAgentConversationsParams{Limit: 10})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, conversations, 1)
|
||||||
|
latestRun, err := al2.queries.GetLatestAgentRunByConversationID(ctx, memsqlc.GetLatestAgentRunByConversationIDParams{ConversationID: secondConversationID})
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, secondConversationID, latestRun.ConversationID)
|
||||||
|
assert.Equal(t, sessionKey, al2.state.GetLastSessionKey())
|
||||||
|
boundRaw, err := al2.memDelegate.GetKV(ctx, pkgroot.NAME, conversationBindingKey(sessionKey))
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, secondConversationID.String(), boundRaw)
|
||||||
|
al2.Stop()
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentLoop_HeartbeatUsesUniqueSessionAndLastPersistedSessionContext(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
al := newCheckpointTestAgentLoop(t, "heartbeat reply")
|
||||||
|
ctx, cancel := context.WithTimeout(t.Context(), 10*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
userSessionKey := "heartbeat-source-session"
|
||||||
|
al.sessions.SetSummary(userSessionKey, "summary from persisted session")
|
||||||
|
require.NoError(t, al.state.SetLastSessionKeyForTarget(ctx, "test", "chat1", userSessionKey))
|
||||||
|
|
||||||
|
response1, err := al.ProcessHeartbeat(ctx, "heartbeat prompt", "test", "chat1")
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Contains(t, response1, "heartbeat reply")
|
||||||
|
|
||||||
|
response2, err := al.ProcessHeartbeat(ctx, "heartbeat prompt", "test", "chat1")
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Contains(t, response2, "heartbeat reply")
|
||||||
|
|
||||||
|
conversations, err := al.queries.ListAgentConversations(ctx, memsqlc.ListAgentConversationsParams{Limit: 10})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, conversations, 2)
|
||||||
|
for _, conv := range conversations {
|
||||||
|
require.NotNil(t, conv.Title)
|
||||||
|
assert.Contains(t, *conv.Title, "heartbeat:")
|
||||||
|
assert.NotEqual(t, "heartbeat", *conv.Title)
|
||||||
|
}
|
||||||
|
assert.Equal(t, userSessionKey, al.state.GetLastSessionKey())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentLoop_HeartbeatUsesTargetScopedSessionContext(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
al := newCheckpointTestAgentLoop(t, "heartbeat reply")
|
||||||
|
ctx, cancel := context.WithTimeout(t.Context(), 10*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
al.sessions.SetSummary("session-chat-a", "summary from chat A")
|
||||||
|
al.sessions.SetSummary("session-chat-b", "summary from chat B")
|
||||||
|
require.NoError(t, al.state.SetLastSessionKeyForTarget(ctx, "telegram", "chat-a", "session-chat-a"))
|
||||||
|
require.NoError(t, al.state.SetLastSessionKeyForTarget(ctx, "telegram", "chat-b", "session-chat-b"))
|
||||||
|
|
||||||
|
responseA, err := al.ProcessHeartbeat(ctx, "heartbeat prompt", "telegram", "chat-a")
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Contains(t, responseA, "heartbeat reply")
|
||||||
|
|
||||||
|
responseB, err := al.ProcessHeartbeat(ctx, "heartbeat prompt", "telegram", "chat-b")
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Contains(t, responseB, "heartbeat reply")
|
||||||
|
|
||||||
|
conversations, err := al.queries.ListAgentConversations(ctx, memsqlc.ListAgentConversationsParams{Limit: 20})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, conversations, 2)
|
||||||
|
for _, conv := range conversations {
|
||||||
|
require.NotNil(t, conv.Title)
|
||||||
|
assert.Contains(t, *conv.Title, "heartbeat:")
|
||||||
|
}
|
||||||
|
assert.Equal(t, "session-chat-b", al.state.GetLastSessionKey())
|
||||||
|
assert.Equal(t, "session-chat-a", al.state.GetLastSessionKeyForTarget("telegram", "chat-a"))
|
||||||
|
assert.Equal(t, "session-chat-b", al.state.GetLastSessionKeyForTarget("telegram", "chat-b"))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentLoop_ProcessDirectStreamingPersistsLastSessionKey(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
al := newCheckpointTestAgentLoop(t, "streaming reply")
|
||||||
|
ctx, cancel := context.WithTimeout(t.Context(), 10*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
sessionKey := "streaming-session"
|
||||||
|
response, err := al.ProcessDirectStreaming(ctx, "stream this", sessionKey, "cli", "stream-chat")
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Contains(t, response, "streaming reply")
|
||||||
|
assert.Equal(t, sessionKey, al.state.GetLastSessionKey())
|
||||||
|
}
|
||||||
|
|
||||||
func newCheckpointTestAgentLoop(t *testing.T, response string) *AgentLoop {
|
func newCheckpointTestAgentLoop(t *testing.T, response string) *AgentLoop {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -114,12 +114,16 @@ func (cb *ContextBuilder) SetSessionResolver(sessionKeyFn func() string) {
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cb *ContextBuilder) getIdentity() string {
|
func (cb *ContextBuilder) getIdentity() string {
|
||||||
|
return cb.getIdentityForQuery("")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (cb *ContextBuilder) getIdentityForQuery(query string) string {
|
||||||
now := time.Now().Format("2006-01-02 15:04 (Monday)")
|
now := time.Now().Format("2006-01-02 15:04 (Monday)")
|
||||||
workspacePath, _ := filepath.Abs(filepath.Join(cb.workspace))
|
workspacePath, _ := filepath.Abs(filepath.Join(cb.workspace))
|
||||||
runtime := fmt.Sprintf("%s %s, Go %s", runtime.GOOS, runtime.GOARCH, runtime.Version())
|
runtime := fmt.Sprintf("%s %s, Go %s", runtime.GOOS, runtime.GOARCH, runtime.Version())
|
||||||
|
|
||||||
// Build tools section dynamically
|
// Build tools section dynamically
|
||||||
toolsSection := cb.buildToolsSection()
|
toolsSection := cb.buildToolsSection(query)
|
||||||
|
|
||||||
return fmt.Sprintf(`# dragonscale 🦞
|
return fmt.Sprintf(`# dragonscale 🦞
|
||||||
|
|
||||||
|
|
@ -170,12 +174,26 @@ Your workspace is at: %s
|
||||||
now, runtime, workspacePath, toolsSection)
|
now, runtime, workspacePath, toolsSection)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cb *ContextBuilder) buildToolsSection() string {
|
func (cb *ContextBuilder) buildToolsSection(query string) string {
|
||||||
if cb.tools == nil {
|
if cb.tools == nil {
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
||||||
summaries := cb.tools.GetSummaries()
|
var summaries []string
|
||||||
|
trimmedQuery := strings.TrimSpace(query)
|
||||||
|
if trimmedQuery == "" {
|
||||||
|
summaries = cb.tools.GetSummaries()
|
||||||
|
} else if !isPlanningOnlyPrompt(trimmedQuery) {
|
||||||
|
names := initialPromptToolNames(cb.tools, trimmedQuery)
|
||||||
|
summaries = make([]string, 0, len(names))
|
||||||
|
for _, name := range names {
|
||||||
|
tool, ok := cb.tools.Get(name)
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
summaries = append(summaries, fmt.Sprintf("- `%s` - %s", tool.Name(), tool.Description()))
|
||||||
|
}
|
||||||
|
}
|
||||||
if len(summaries) == 0 {
|
if len(summaries) == 0 {
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
@ -199,16 +217,24 @@ type contextSection struct {
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cb *ContextBuilder) BuildSystemPrompt() string {
|
func (cb *ContextBuilder) BuildSystemPrompt() string {
|
||||||
return cb.BuildSystemPromptWithBudget(cb.tokenBudgetTokens())
|
return cb.BuildSystemPromptForTurn("", "", cb.tokenBudgetTokens())
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cb *ContextBuilder) BuildSystemPromptWithBudget(budgetTokens int) string {
|
func (cb *ContextBuilder) BuildSystemPromptWithBudget(budgetTokens int) string {
|
||||||
|
return cb.BuildSystemPromptForTurn("", "", budgetTokens)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (cb *ContextBuilder) BuildSystemPromptForSession(sessionKey string, budgetTokens int) string {
|
||||||
|
return cb.BuildSystemPromptForTurn(sessionKey, "", budgetTokens)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (cb *ContextBuilder) BuildSystemPromptForTurn(sessionKey, query string, budgetTokens int) string {
|
||||||
|
|
||||||
// Collect sections in priority order
|
// Collect sections in priority order
|
||||||
sections := []contextSection{}
|
sections := []contextSection{}
|
||||||
|
|
||||||
// P0: Core identity (always included)
|
// P0: Core identity (always included)
|
||||||
sections = append(sections, contextSection{"identity", cb.getIdentity(), 0})
|
sections = append(sections, contextSection{"identity", cb.getIdentityForQuery(query), 0})
|
||||||
|
|
||||||
// P1: Bootstrap files (user identity) — cached with TTL
|
// P1: Bootstrap files (user identity) — cached with TTL
|
||||||
if bc := cb.cachedBootstrapFiles(); bc != "" {
|
if bc := cb.cachedBootstrapFiles(); bc != "" {
|
||||||
|
|
@ -231,7 +257,7 @@ Do NOT assume skill content — always load before applying.
|
||||||
|
|
||||||
// P3: Working context (hot tier — highly dynamic, high value)
|
// P3: Working context (hot tier — highly dynamic, high value)
|
||||||
if cb.memoryStore != nil {
|
if cb.memoryStore != nil {
|
||||||
if wc := cb.buildWorkingContextSection(); wc != "" {
|
if wc := cb.buildWorkingContextSection(sessionKey); wc != "" {
|
||||||
sections = append(sections, contextSection{"working_context", wc, 3})
|
sections = append(sections, contextSection{"working_context", wc, 3})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -509,7 +535,7 @@ func (cb *ContextBuilder) LoadBootstrapFiles() string {
|
||||||
|
|
||||||
// buildWorkingContextSection returns the working context section for the system prompt.
|
// buildWorkingContextSection returns the working context section for the system prompt.
|
||||||
// It includes the hot-tier working context buffer and memory usage instructions.
|
// It includes the hot-tier working context buffer and memory usage instructions.
|
||||||
func (cb *ContextBuilder) buildWorkingContextSection() string {
|
func (cb *ContextBuilder) buildWorkingContextSection(sessionKey string) string {
|
||||||
if cb.memoryStore == nil {
|
if cb.memoryStore == nil {
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
@ -520,15 +546,18 @@ func (cb *ContextBuilder) buildWorkingContextSection() string {
|
||||||
|
|
||||||
var parts []string
|
var parts []string
|
||||||
|
|
||||||
sessionKey := "default"
|
resolvedSessionKey := strings.TrimSpace(sessionKey)
|
||||||
if cb.sessionKeyFn != nil {
|
if resolvedSessionKey == "" && cb.sessionKeyFn != nil {
|
||||||
if resolved := strings.TrimSpace(cb.sessionKeyFn()); resolved != "" {
|
if resolved := strings.TrimSpace(cb.sessionKeyFn()); resolved != "" {
|
||||||
sessionKey = resolved
|
resolvedSessionKey = resolved
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
if resolvedSessionKey == "" {
|
||||||
|
resolvedSessionKey = "default"
|
||||||
|
}
|
||||||
|
|
||||||
// Inject working context (hot tier)
|
// Inject working context (hot tier)
|
||||||
wc, err := cb.memoryStore.GetWorkingContext(ctx, pkg.NAME, sessionKey)
|
wc, err := cb.memoryStore.GetWorkingContext(ctx, pkg.NAME, resolvedSessionKey)
|
||||||
if err == nil && wc != "" {
|
if err == nil && wc != "" {
|
||||||
parts = append(parts, "## Working Context\n\n"+wc)
|
parts = append(parts, "## Working Context\n\n"+wc)
|
||||||
}
|
}
|
||||||
|
|
@ -552,10 +581,10 @@ Store important user preferences, key decisions, and facts you want to remember
|
||||||
return strings.Join(parts, "\n\n")
|
return strings.Join(parts, "\n\n")
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cb *ContextBuilder) BuildMessages(history []messages.Message, summary string, currentMessage string, media []string, channel, chatID string) []messages.Message {
|
func (cb *ContextBuilder) BuildMessages(sessionKey string, history []messages.Message, summary string, currentMessage string, media []string, channel, chatID string) []messages.Message {
|
||||||
msgs := []messages.Message{}
|
msgs := []messages.Message{}
|
||||||
|
|
||||||
systemPrompt := cb.BuildSystemPrompt()
|
systemPrompt := cb.BuildSystemPromptForTurn(sessionKey, currentMessage, cb.tokenBudgetTokens())
|
||||||
|
|
||||||
// Add Current Session info if provided
|
// Add Current Session info if provided
|
||||||
if channel != "" && chatID != "" {
|
if channel != "" && chatID != "" {
|
||||||
|
|
|
||||||
|
|
@ -3,6 +3,8 @@ package agent
|
||||||
import (
|
import (
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
|
"github.com/ZanzyTHEbar/dragonscale/pkg/tools"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestSystemPromptIncludesDirectToolRoutingHints(t *testing.T) {
|
func TestSystemPromptIncludesDirectToolRoutingHints(t *testing.T) {
|
||||||
|
|
@ -24,3 +26,25 @@ func TestSystemPromptIncludesDirectToolRoutingHints(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestSystemPromptForTurnLimitsToolSectionToRelevantHints(t *testing.T) {
|
||||||
|
cb := NewContextBuilder(t.TempDir())
|
||||||
|
reg := tools.NewToolRegistry()
|
||||||
|
reg.Register(&namedTool{name: "memory"})
|
||||||
|
reg.Register(&namedTool{name: "obligation"})
|
||||||
|
reg.Register(&namedTool{name: "write_file"})
|
||||||
|
reg.Register(&namedTool{name: "read_file"})
|
||||||
|
cb.SetToolsRegistry(reg)
|
||||||
|
|
||||||
|
prompt := cb.BuildSystemPromptForTurn("session-a", "Capture these commitments and give me a reminder/follow-up plan with explicit timing.", 0)
|
||||||
|
|
||||||
|
if !strings.Contains(prompt, "`memory`") {
|
||||||
|
t.Fatalf("expected turn-specific prompt to include memory tool summary, got: %s", prompt)
|
||||||
|
}
|
||||||
|
if strings.Contains(prompt, "`obligation`") {
|
||||||
|
t.Fatalf("did not expect turn-specific prompt to advertise obligation, got: %s", prompt)
|
||||||
|
}
|
||||||
|
if strings.Contains(prompt, "`write_file`") || strings.Contains(prompt, "`read_file`") {
|
||||||
|
t.Fatalf("did not expect unrelated file tools in turn-specific prompt, got: %s", prompt)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -166,6 +166,40 @@ func TestGroundFinalContentExpandsExactCommitmentRegister(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestGroundFinalContentExpandsDailyPlanWeekdayAbbreviations(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
al := &AgentLoop{}
|
||||||
|
got := al.groundFinalContent(
|
||||||
|
"Given prior commitments {invoice Monday, PR review Tuesday, dentist this month}, provide this week's daily plan and explicitly carry forward unfinished items.",
|
||||||
|
"**Week of April 20–26, 2026**\n\n- **Mon 4/20** — Submit invoice\n- **Tue 4/21** — Complete PR review\n- **Fri 4/24** — Dentist appointment\n\n**Carry-forward if unfinished**\n- Invoice -> Tue\n- PR review -> Wed\n- Dentist -> next available weekday",
|
||||||
|
nil,
|
||||||
|
)
|
||||||
|
|
||||||
|
for _, snippet := range []string{"Monday 4/20", "Tuesday 4/21", "Friday 4/24"} {
|
||||||
|
if !strings.Contains(got, snippet) {
|
||||||
|
t.Fatalf("expected grounded continuity plan to include %q, got %q", snippet, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGroundFinalContentAddsWebinarFollowUpLanguage(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
al := &AgentLoop{}
|
||||||
|
got := al.groundFinalContent(
|
||||||
|
"I need to launch a small webinar next week. Give me a plan that includes proactive risk checks and follow-up actions I might forget.",
|
||||||
|
"Pre-launch: risk check on internet backup and dial-in fallback.\n\nLaunch day: 30 min early.\n\nPost-event: send recording within 24 hours.",
|
||||||
|
nil,
|
||||||
|
)
|
||||||
|
|
||||||
|
for _, snippet := range []string{"Follow-up actions", "follow-up check-in", "verify attendee follow-up status"} {
|
||||||
|
if !strings.Contains(strings.ToLower(got), strings.ToLower(snippet)) {
|
||||||
|
t.Fatalf("expected grounded webinar plan to include %q, got %q", snippet, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestGroundFinalContentRecoversSkillSummary(t *testing.T) {
|
func TestGroundFinalContentRecoversSkillSummary(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
|
|
@ -237,6 +271,52 @@ func TestResolveFinalContentPrefersToolResultOverPreamble(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestResolveFinalContentFallsBackToClarificationWhenNoTextRecovered(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
al := &AgentLoop{}
|
||||||
|
got, err := al.resolveFinalContent("", []fantasy.StepResult{
|
||||||
|
stepWithTextAndToolResults("Which file do you mean, and what do you want me to do with it?"),
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("resolveFinalContent returned error: %v", err)
|
||||||
|
}
|
||||||
|
if !strings.Contains(strings.ToLower(got), "what do you want me to do") {
|
||||||
|
t.Fatalf("expected clarification fallback, got %q", got)
|
||||||
|
}
|
||||||
|
if !strings.Contains(strings.ToLower(got), "which file do you mean") {
|
||||||
|
t.Fatalf("expected clarification fallback to preserve clarification text, got %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveFinalContentReturnsErrorForEmptyClearNoToolRun(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
al := &AgentLoop{}
|
||||||
|
_, err := al.resolveFinalContent("", []fantasy.StepResult{{}})
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("expected clear empty no-tool run to return an error")
|
||||||
|
}
|
||||||
|
if !strings.Contains(err.Error(), "no final response text") {
|
||||||
|
t.Fatalf("expected empty-response error, got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveFinalContentIgnoresLowConfidenceNoToolPreamble(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
al := &AgentLoop{}
|
||||||
|
_, err := al.resolveFinalContent("", []fantasy.StepResult{
|
||||||
|
stepWithTextAndToolResults("Let me think about that for a moment:"),
|
||||||
|
})
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("expected low-confidence no-tool preamble to still return an error")
|
||||||
|
}
|
||||||
|
if !strings.Contains(err.Error(), "no final response text") {
|
||||||
|
t.Fatalf("expected empty-response error, got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestGroundFinalContentOverridesContradictoryExecSuccess(t *testing.T) {
|
func TestGroundFinalContentOverridesContradictoryExecSuccess(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
|
|
@ -253,6 +333,22 @@ func TestGroundFinalContentOverridesContradictoryExecSuccess(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestGroundFinalContentOverridesContradictoryExecDenial(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
al := &AgentLoop{}
|
||||||
|
got := al.groundFinalContent(
|
||||||
|
"Run the command 'echo progressive-test-marker' and tell me the output.",
|
||||||
|
"The command execution was denied. I don't have permission to run shell commands in this environment.",
|
||||||
|
[]fantasy.StepResult{
|
||||||
|
stepWithToolResults(toolText("exec", "progressive-test-marker")),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
if !strings.Contains(got, "progressive-test-marker") {
|
||||||
|
t.Fatalf("expected grounded exec success output, got %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func stepWithToolResults(results ...fantasy.ToolResultContent) fantasy.StepResult {
|
func stepWithToolResults(results ...fantasy.ToolResultContent) fantasy.StepResult {
|
||||||
content := make(fantasy.ResponseContent, 0, len(results))
|
content := make(fantasy.ResponseContent, 0, len(results))
|
||||||
for _, result := range results {
|
for _, result := range results {
|
||||||
|
|
|
||||||
|
|
@ -2,6 +2,7 @@ package agent
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/ZanzyTHEbar/dragonscale/pkg/skills"
|
"github.com/ZanzyTHEbar/dragonscale/pkg/skills"
|
||||||
|
|
@ -62,13 +63,13 @@ func TestInitialPromptToolsExposeSkillAndFileHelpers(t *testing.T) {
|
||||||
}
|
}
|
||||||
|
|
||||||
spawnNames := toolNames(al.initialPromptTools("Spawn a background task to write the text 'async-spawn-test' to a file called spawn_output.txt."))
|
spawnNames := toolNames(al.initialPromptTools("Spawn a background task to write the text 'async-spawn-test' to a file called spawn_output.txt."))
|
||||||
if !containsAll(spawnNames, "spawn", "write_file") {
|
if len(spawnNames) != 1 || !containsAll(spawnNames, "spawn") {
|
||||||
t.Fatalf("expected spawn/write_file, got %v", spawnNames)
|
t.Fatalf("expected delegation prompt to expose only spawn, got %v", spawnNames)
|
||||||
}
|
}
|
||||||
|
|
||||||
subagentNames := toolNames(al.initialPromptTools("Use a subagent to calculate the sum of 10 + 20 + 30 and report the result back to me."))
|
subagentNames := toolNames(al.initialPromptTools("Use a subagent to calculate the sum of 10 + 20 + 30 and report the result back to me."))
|
||||||
if !containsAll(subagentNames, "subagent") {
|
if len(subagentNames) != 1 || !containsAll(subagentNames, "subagent") {
|
||||||
t.Fatalf("expected subagent, got %v", subagentNames)
|
t.Fatalf("expected delegation prompt to expose only subagent, got %v", subagentNames)
|
||||||
}
|
}
|
||||||
|
|
||||||
memoryNames := toolNames(al.initialPromptTools("Track these commitments exactly: send rent receipt tonight, book vet appointment tomorrow, and submit sprint notes by Friday."))
|
memoryNames := toolNames(al.initialPromptTools("Track these commitments exactly: send rent receipt tonight, book vet appointment tomorrow, and submit sprint notes by Friday."))
|
||||||
|
|
@ -76,6 +77,14 @@ func TestInitialPromptToolsExposeSkillAndFileHelpers(t *testing.T) {
|
||||||
t.Fatalf("expected memory, got %v", memoryNames)
|
t.Fatalf("expected memory, got %v", memoryNames)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
reminderNames := toolNames(al.initialPromptTools("Schedule reminder to pay rent tomorrow at 9am."))
|
||||||
|
if !containsAll(reminderNames, "obligation") {
|
||||||
|
t.Fatalf("expected obligation for explicit reminder scheduling, got %v", reminderNames)
|
||||||
|
}
|
||||||
|
if isPlanningOnlyPrompt("Schedule reminder to pay rent tomorrow at 9am.") {
|
||||||
|
t.Fatal("expected explicit reminder scheduling prompt to stay actionable")
|
||||||
|
}
|
||||||
|
|
||||||
memorySearchNames := toolNames(al.initialPromptTools("Search your memory for 'xyzzy_nonexistent_topic_42' and tell me what you find."))
|
memorySearchNames := toolNames(al.initialPromptTools("Search your memory for 'xyzzy_nonexistent_topic_42' and tell me what you find."))
|
||||||
if !containsAll(memorySearchNames, "memory") {
|
if !containsAll(memorySearchNames, "memory") {
|
||||||
t.Fatalf("expected memory for memory-search prompt, got %v", memorySearchNames)
|
t.Fatalf("expected memory for memory-search prompt, got %v", memorySearchNames)
|
||||||
|
|
@ -107,11 +116,159 @@ func TestIsPlanningOnlyPrompt(t *testing.T) {
|
||||||
t.Fatal("expected capture/reminder prompt not to be treated as planning-only")
|
t.Fatal("expected capture/reminder prompt not to be treated as planning-only")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if isPlanningOnlyPrompt("Schedule reminder to pay rent tomorrow at 9am.") {
|
||||||
|
t.Fatal("expected explicit reminder scheduling prompt not to be treated as planning-only")
|
||||||
|
}
|
||||||
|
|
||||||
if !isPlanningOnlyPrompt("I must send a proposal in 4 hours. Give me a reminder schedule and specify when the first reminder should fire.") {
|
if !isPlanningOnlyPrompt("I must send a proposal in 4 hours. Give me a reminder schedule and specify when the first reminder should fire.") {
|
||||||
t.Fatal("expected reminder schedule request to be treated as planning-only")
|
t.Fatal("expected reminder schedule request to be treated as planning-only")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestTurnConstraintForQuery_DelegationFirst(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
constraint := turnConstraintForQuery("Use a subagent to calculate the sum of 10 + 20 + 30 and report the result back to me.")
|
||||||
|
if !strings.Contains(constraint, "Call `subagent` as your first tool step") {
|
||||||
|
t.Fatalf("expected subagent-first constraint, got %q", constraint)
|
||||||
|
}
|
||||||
|
if !strings.Contains(constraint, "Do not use tool_search or tool_call first") {
|
||||||
|
t.Fatalf("expected delegation constraint to forbid tool_search/tool_call detours, got %q", constraint)
|
||||||
|
}
|
||||||
|
|
||||||
|
spawnConstraint := turnConstraintForQuery("Spawn a background task to write the text 'async-spawn-test' to a file called spawn_output.txt.")
|
||||||
|
if !strings.Contains(spawnConstraint, "Call `spawn` as your first tool step") {
|
||||||
|
t.Fatalf("expected spawn-first constraint, got %q", spawnConstraint)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTurnConstraintForQuery_DoesNotForceDelegationForGenericAsyncText(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
for _, query := range []string{
|
||||||
|
"Debug this async callback regression and explain the root cause.",
|
||||||
|
"Review async code paths and suggest fixes.",
|
||||||
|
"Explain async behavior in this runtime.",
|
||||||
|
} {
|
||||||
|
if constraint := turnConstraintForQuery(query); strings.Contains(constraint, "Call `spawn` as your first tool step") {
|
||||||
|
t.Fatalf("expected generic async prompt not to force spawn delegation, got %q for %q", constraint, query)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTurnConstraintForQuery_DoesNotForceDelegationForMetaDelegationPrompts(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
for _, query := range []string{
|
||||||
|
"When should we delegate this task to a subagent?",
|
||||||
|
"Give me a plan to delegate this work safely.",
|
||||||
|
"Explain whether we should use a subagent here.",
|
||||||
|
"Why is this running in the background?",
|
||||||
|
"Explain how this runs asynchronously.",
|
||||||
|
"Use a subagent or handle it directly?",
|
||||||
|
"Run this in the background?",
|
||||||
|
} {
|
||||||
|
constraint := turnConstraintForQuery(query)
|
||||||
|
if strings.Contains(constraint, "Call `spawn` as your first tool step") || strings.Contains(constraint, "Call `subagent` as your first tool step") {
|
||||||
|
t.Fatalf("expected meta/delegation discussion prompt not to force delegation, got %q for %q", constraint, query)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestInitialPromptTools_DoNotExposeSpawnForGenericAsyncText(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
reg := tools.NewToolRegistry()
|
||||||
|
reg.Register(&namedTool{name: "tool_search"})
|
||||||
|
reg.Register(&namedTool{name: "tool_call"})
|
||||||
|
reg.Register(&namedTool{name: "spawn"})
|
||||||
|
reg.Register(&namedTool{name: "subagent"})
|
||||||
|
|
||||||
|
al := &AgentLoop{tools: reg}
|
||||||
|
|
||||||
|
for _, query := range []string{
|
||||||
|
"Debug this async callback regression and explain the root cause.",
|
||||||
|
"Review async code paths and suggest fixes.",
|
||||||
|
} {
|
||||||
|
names := toolNames(al.initialPromptTools(query))
|
||||||
|
if containsAll(names, "spawn") {
|
||||||
|
t.Fatalf("expected generic async prompt not to expose spawn, got %v for %q", names, query)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestInitialPromptTools_DoNotExposeDelegationToolsForMetaDelegationPrompts(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
reg := tools.NewToolRegistry()
|
||||||
|
reg.Register(&namedTool{name: "tool_search"})
|
||||||
|
reg.Register(&namedTool{name: "tool_call"})
|
||||||
|
reg.Register(&namedTool{name: "spawn"})
|
||||||
|
reg.Register(&namedTool{name: "subagent"})
|
||||||
|
|
||||||
|
al := &AgentLoop{tools: reg}
|
||||||
|
|
||||||
|
for _, query := range []string{
|
||||||
|
"When should we delegate this task to a subagent?",
|
||||||
|
"Give me a plan to delegate this work safely.",
|
||||||
|
"Explain whether we should use a subagent here.",
|
||||||
|
"Why is this running in the background?",
|
||||||
|
"Use a subagent or handle it directly?",
|
||||||
|
} {
|
||||||
|
names := toolNames(al.initialPromptTools(query))
|
||||||
|
if containsAll(names, "spawn") || containsAll(names, "subagent") || containsAll(names, "tool_search") {
|
||||||
|
t.Fatalf("expected meta/delegation discussion prompt not to expose delegation tools, got %v for %q", names, query)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestInitialPromptTools_DefaultToolSearchOnlyForActionableOpenEndedPrompts(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
reg := tools.NewToolRegistry()
|
||||||
|
reg.Register(&namedTool{name: "tool_search"})
|
||||||
|
reg.Register(&namedTool{name: "tool_call"})
|
||||||
|
|
||||||
|
al := &AgentLoop{tools: reg}
|
||||||
|
|
||||||
|
actionable := toolNames(al.initialPromptTools("Debug this flaky worker startup issue."))
|
||||||
|
if len(actionable) != 1 || !containsAll(actionable, "tool_search") {
|
||||||
|
t.Fatalf("expected actionable open-ended prompt to expose tool_search, got %v", actionable)
|
||||||
|
}
|
||||||
|
|
||||||
|
meta := toolNames(al.initialPromptTools("Why is this running in the background?"))
|
||||||
|
if len(meta) != 0 {
|
||||||
|
t.Fatalf("expected meta discussion prompt not to expose tool_search, got %v", meta)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTurnConstraintForQuery_PlanningOnlyCompactResponse(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
constraint := turnConstraintForQuery("Given prior commitments {invoice Monday, PR review Tuesday, dentist this month}, provide this week's daily plan and explicitly carry forward unfinished items.")
|
||||||
|
if !strings.Contains(constraint, "planning-only") {
|
||||||
|
t.Fatalf("expected planning-only constraint, got %q", constraint)
|
||||||
|
}
|
||||||
|
if !strings.Contains(constraint, "Keep the answer compact and structured") {
|
||||||
|
t.Fatalf("expected compact-response constraint, got %q", constraint)
|
||||||
|
}
|
||||||
|
if !strings.Contains(constraint, "brief day/week bullets") {
|
||||||
|
t.Fatalf("expected structured brevity guidance, got %q", constraint)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExplicitWriteFileRequest_HandlesQuotedSpecialFilename(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
path, content := explicitWriteFileRequest("Create a file called 'test file (1).txt' with the content 'special chars test' and confirm success.")
|
||||||
|
if path != "test file (1).txt" {
|
||||||
|
t.Fatalf("expected quoted filename to be extracted, got %q", path)
|
||||||
|
}
|
||||||
|
if content != "special chars test" {
|
||||||
|
t.Fatalf("expected quoted content to be extracted, got %q", content)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func containsAll(have []string, want ...string) bool {
|
func containsAll(have []string, want ...string) bool {
|
||||||
set := make(map[string]struct{}, len(have))
|
set := make(map[string]struct{}, len(have))
|
||||||
for _, name := range have {
|
for _, name := range have {
|
||||||
|
|
|
||||||
|
|
@ -5,6 +5,7 @@ import (
|
||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
"strings"
|
"strings"
|
||||||
|
"sync"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
|
@ -157,6 +158,160 @@ func (t *echoTool) Execute(_ context.Context, args map[string]interface{}) *tool
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type spawnLikeAsyncTool struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
channel string
|
||||||
|
chatID string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *spawnLikeAsyncTool) Name() string { return "spawn" }
|
||||||
|
func (t *spawnLikeAsyncTool) Description() string { return "Test async spawn-like tool" }
|
||||||
|
func (t *spawnLikeAsyncTool) Parameters() map[string]interface{} {
|
||||||
|
return map[string]interface{}{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]interface{}{
|
||||||
|
"task": map[string]interface{}{
|
||||||
|
"type": "string",
|
||||||
|
"description": "Task description",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": []string{"task"},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *spawnLikeAsyncTool) SetContext(channel, chatID string) {
|
||||||
|
t.mu.Lock()
|
||||||
|
defer t.mu.Unlock()
|
||||||
|
t.channel = channel
|
||||||
|
t.chatID = chatID
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *spawnLikeAsyncTool) SetCallback(cb tools.AsyncCallback) {
|
||||||
|
_ = cb
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *spawnLikeAsyncTool) Execute(ctx context.Context, _ map[string]interface{}) *tools.ToolResult {
|
||||||
|
channel, chatID := tools.ResolveExecutionTarget(ctx, "", "")
|
||||||
|
callback := tools.AsyncCallbackFromContext(ctx)
|
||||||
|
t.mu.Lock()
|
||||||
|
t.channel = channel
|
||||||
|
t.chatID = chatID
|
||||||
|
t.mu.Unlock()
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
if callback != nil {
|
||||||
|
callback(ctx, &tools.ToolResult{
|
||||||
|
ForLLM: "background task finished",
|
||||||
|
ForUser: fmt.Sprintf("async completion on %s:%s", channel, chatID),
|
||||||
|
Silent: false,
|
||||||
|
IsError: false,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
return tools.AsyncResult("background task started")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *spawnLikeAsyncTool) Context() (string, string) {
|
||||||
|
t.mu.Lock()
|
||||||
|
defer t.mu.Unlock()
|
||||||
|
return t.channel, t.chatID
|
||||||
|
}
|
||||||
|
|
||||||
|
type spawnToolCallingModel struct {
|
||||||
|
callCount int
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *spawnToolCallingModel) Generate(_ context.Context, call fantasy.Call) (*fantasy.Response, error) {
|
||||||
|
m.callCount++
|
||||||
|
|
||||||
|
hasToolResults := false
|
||||||
|
for _, msg := range call.Prompt {
|
||||||
|
for _, part := range msg.Content {
|
||||||
|
if part.GetType() == fantasy.ContentTypeToolResult {
|
||||||
|
hasToolResults = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if !hasToolResults {
|
||||||
|
for _, tool := range call.Tools {
|
||||||
|
if tool.GetName() == "spawn" {
|
||||||
|
return &fantasy.Response{
|
||||||
|
Content: fantasy.ResponseContent{
|
||||||
|
fantasy.ToolCallContent{
|
||||||
|
ToolCallID: "call-spawn-1",
|
||||||
|
ToolName: "spawn",
|
||||||
|
Input: `{"task":"background work"}`,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
FinishReason: fantasy.FinishReasonToolCalls,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return &fantasy.Response{
|
||||||
|
Content: fantasy.ResponseContent{
|
||||||
|
fantasy.TextContent{Text: "Spawn request completed"},
|
||||||
|
},
|
||||||
|
FinishReason: fantasy.FinishReasonStop,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *spawnToolCallingModel) Stream(ctx context.Context, call fantasy.Call) (fantasy.StreamResponse, error) {
|
||||||
|
resp, err := m.Generate(ctx, call)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return func(yield func(fantasy.StreamPart) bool) {
|
||||||
|
hasToolCalls := false
|
||||||
|
for _, c := range resp.Content {
|
||||||
|
if c.GetType() == fantasy.ContentTypeToolCall {
|
||||||
|
hasToolCalls = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if hasToolCalls {
|
||||||
|
for _, c := range resp.Content {
|
||||||
|
if tc, ok := c.(fantasy.ToolCallContent); ok {
|
||||||
|
if !yield(fantasy.StreamPart{Type: fantasy.StreamPartTypeToolCall, ID: tc.ToolCallID, ToolCallName: tc.ToolName, ToolCallInput: tc.Input}) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
yield(fantasy.StreamPart{Type: fantasy.StreamPartTypeFinish, FinishReason: fantasy.FinishReasonToolCalls})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
text := resp.Content.Text()
|
||||||
|
if text != "" {
|
||||||
|
if !yield(fantasy.StreamPart{Type: fantasy.StreamPartTypeTextStart, ID: "text-0"}) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if !yield(fantasy.StreamPart{Type: fantasy.StreamPartTypeTextDelta, ID: "text-0", Delta: text}) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if !yield(fantasy.StreamPart{Type: fantasy.StreamPartTypeTextEnd, ID: "text-0"}) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
yield(fantasy.StreamPart{Type: fantasy.StreamPartTypeFinish, FinishReason: fantasy.FinishReasonStop})
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *spawnToolCallingModel) GenerateObject(_ context.Context, _ fantasy.ObjectCall) (*fantasy.ObjectResponse, error) {
|
||||||
|
return nil, fmt.Errorf("not implemented")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *spawnToolCallingModel) StreamObject(_ context.Context, _ fantasy.ObjectCall) (fantasy.ObjectStreamResponse, error) {
|
||||||
|
return nil, fmt.Errorf("not implemented")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *spawnToolCallingModel) Provider() string { return "mock" }
|
||||||
|
func (m *spawnToolCallingModel) Model() string { return "mock-spawn-model" }
|
||||||
|
|
||||||
// --- Integration Tests ---
|
// --- Integration Tests ---
|
||||||
|
|
||||||
// TestIntegration_FullAgentLoop_SimpleResponse tests the full agent loop
|
// TestIntegration_FullAgentLoop_SimpleResponse tests the full agent loop
|
||||||
|
|
@ -278,6 +433,78 @@ func TestIntegration_FullAgentLoop_WithToolCalls(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestIntegration_SecureBusExecutor_PreservesAsyncCallbacksAndContext(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
tmpDir, err := os.MkdirTemp("", "agent-integration-async-*")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to create temp dir: %v", err)
|
||||||
|
}
|
||||||
|
defer os.RemoveAll(tmpDir)
|
||||||
|
|
||||||
|
cfg := &config.Config{
|
||||||
|
Agents: config.AgentsConfig{
|
||||||
|
Defaults: config.AgentDefaults{
|
||||||
|
Sandbox: tmpDir,
|
||||||
|
Model: "mock-spawn-model",
|
||||||
|
MaxTokens: 4096,
|
||||||
|
MaxToolIterations: 10,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
msgBus := bus.NewMessageBus()
|
||||||
|
model := &spawnToolCallingModel{}
|
||||||
|
al := mustNewAgentLoop(t, cfg, msgBus, model)
|
||||||
|
spawnTool := &spawnLikeAsyncTool{}
|
||||||
|
al.RegisterTool(spawnTool)
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(t.Context(), 10*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
outboundCh := make(chan bus.OutboundMessage, 1)
|
||||||
|
go func() {
|
||||||
|
msg, ok := msgBus.SubscribeOutbound(ctx)
|
||||||
|
if ok {
|
||||||
|
outboundCh <- msg
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
response, err := al.processMessage(ctx, bus.InboundMessage{
|
||||||
|
Channel: "telegram",
|
||||||
|
SenderID: "user1",
|
||||||
|
ChatID: "chat-123",
|
||||||
|
Content: "Spawn a background task to do some work.",
|
||||||
|
SessionKey: "async-session",
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("processMessage failed: %v", err)
|
||||||
|
}
|
||||||
|
if !strings.Contains(response, "Spawn request completed") {
|
||||||
|
t.Fatalf("expected final response after async tool call, got: %s", response)
|
||||||
|
}
|
||||||
|
|
||||||
|
select {
|
||||||
|
case outbound := <-outboundCh:
|
||||||
|
if outbound.Channel != "telegram" || outbound.ChatID != "chat-123" {
|
||||||
|
t.Fatalf("unexpected outbound target: %s:%s", outbound.Channel, outbound.ChatID)
|
||||||
|
}
|
||||||
|
if outbound.Content != "async completion on telegram:chat-123" {
|
||||||
|
t.Fatalf("unexpected async completion message: %s", outbound.Content)
|
||||||
|
}
|
||||||
|
case <-ctx.Done():
|
||||||
|
t.Fatal("timed out waiting for async completion outbound message")
|
||||||
|
}
|
||||||
|
|
||||||
|
channel, chatID := spawnTool.Context()
|
||||||
|
if channel != "telegram" || chatID != "chat-123" {
|
||||||
|
t.Fatalf("context was not propagated to tool: %s:%s", channel, chatID)
|
||||||
|
}
|
||||||
|
if model.callCount < 2 {
|
||||||
|
t.Fatalf("expected at least 2 model calls, got %d", model.callCount)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// TestIntegration_ProcessDirect tests the ProcessDirect method
|
// TestIntegration_ProcessDirect tests the ProcessDirect method
|
||||||
// which is used by CLI mode for one-shot message processing.
|
// which is used by CLI mode for one-shot message processing.
|
||||||
func TestIntegration_ProcessDirect(t *testing.T) {
|
func TestIntegration_ProcessDirect(t *testing.T) {
|
||||||
|
|
|
||||||
|
|
@ -348,9 +348,9 @@ func NewAgentLoop(ctx context.Context, cfg *config.Config, msgBus *bus.MessageBu
|
||||||
contextBuilder.SetSessionResolver(sessionKeyFn)
|
contextBuilder.SetSessionResolver(sessionKeyFn)
|
||||||
memTool.SetSessionResolver(sessionKeyFn)
|
memTool.SetSessionResolver(sessionKeyFn)
|
||||||
subagentMemTool.SetSessionResolver(sessionKeyFn)
|
subagentMemTool.SetSessionResolver(sessionKeyFn)
|
||||||
focusInvalidate := func() {
|
focusInvalidate := func(sessionKey string) {
|
||||||
if sk := sessionKeyFn(); sk != "" {
|
if sessionKey != "" {
|
||||||
al.focusDirty.Store(sk, struct{}{})
|
al.focusDirty.Store(sessionKey, struct{}{})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
startFocus := tools.NewStartFocusTool(memDelegate, sessionsManager, sessionKeyFn)
|
startFocus := tools.NewStartFocusTool(memDelegate, sessionsManager, sessionKeyFn)
|
||||||
|
|
@ -459,9 +459,11 @@ func (al *AgentLoop) Run(ctx context.Context) error {
|
||||||
}
|
}
|
||||||
|
|
||||||
al.inflight.Add(1)
|
al.inflight.Add(1)
|
||||||
|
roundTracker := tools.NewMessageSendTracker()
|
||||||
|
roundCtx := tools.WithMessageSendTracker(ctx, roundTracker)
|
||||||
response, err := func() (string, error) {
|
response, err := func() (string, error) {
|
||||||
defer al.inflight.Done()
|
defer al.inflight.Done()
|
||||||
return al.processMessage(ctx, msg)
|
return al.processMessage(roundCtx, msg)
|
||||||
}()
|
}()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
response = fmt.Sprintf("Error processing message: %v", err)
|
response = fmt.Sprintf("Error processing message: %v", err)
|
||||||
|
|
@ -470,12 +472,7 @@ func (al *AgentLoop) Run(ctx context.Context) error {
|
||||||
if response != "" {
|
if response != "" {
|
||||||
// Check if the message tool already sent a response during this round.
|
// Check if the message tool already sent a response during this round.
|
||||||
// If so, skip publishing to avoid duplicate messages to the user.
|
// If so, skip publishing to avoid duplicate messages to the user.
|
||||||
alreadySent := false
|
alreadySent := roundTracker.Sent()
|
||||||
if tool, ok := al.tools.Get("message"); ok {
|
|
||||||
if mt, ok := tool.(*tools.MessageTool); ok {
|
|
||||||
alreadySent = mt.HasSentInRound()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if !alreadySent {
|
if !alreadySent {
|
||||||
outMsg := bus.OutboundMessage{
|
outMsg := bus.OutboundMessage{
|
||||||
|
|
@ -596,7 +593,24 @@ func (al *AgentLoop) SetupSecureBus(ss *security.SecretStore, cfg securebus.BusC
|
||||||
return tools.ExtractCapabilities(t), true
|
return tools.ExtractCapabilities(t), true
|
||||||
}
|
}
|
||||||
executor := func(ctx context.Context, name string, args map[string]interface{}) *tools.ToolResult {
|
executor := func(ctx context.Context, name string, args map[string]interface{}) *tools.ToolResult {
|
||||||
return al.tools.Execute(ctx, name, args)
|
if sessionKey := toolSessionKeyFromContext(ctx); sessionKey != "" {
|
||||||
|
ctx = tools.WithSessionKey(ctx, sessionKey)
|
||||||
|
}
|
||||||
|
channel, chatID := tools.ExecutionTargetFromContext(ctx)
|
||||||
|
var asyncCallback tools.AsyncCallback
|
||||||
|
if al.bus != nil && channel != "" && chatID != "" {
|
||||||
|
asyncCallback = func(_ context.Context, result *tools.ToolResult) {
|
||||||
|
if result == nil || result.ForUser == "" || result.Silent {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
al.bus.PublishOutbound(bus.OutboundMessage{
|
||||||
|
Channel: channel,
|
||||||
|
ChatID: chatID,
|
||||||
|
Content: result.ForUser,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return al.tools.ExecuteWithContext(ctx, name, args, channel, chatID, asyncCallback)
|
||||||
}
|
}
|
||||||
b := securebus.New(cfg, ss, capLookup, executor)
|
b := securebus.New(cfg, ss, capLookup, executor)
|
||||||
al.secureBus = b
|
al.secureBus = b
|
||||||
|
|
|
||||||
|
|
@ -96,7 +96,7 @@ func (t *MemGPTTool) Execute(ctx context.Context, args map[string]interface{}) *
|
||||||
return tools.ErrorResult("invalid arguments: " + err.Error())
|
return tools.ErrorResult("invalid arguments: " + err.Error())
|
||||||
}
|
}
|
||||||
|
|
||||||
result, err := memstore.NewMemoryTool(t.store, t.agentID, t.currentSession()).Execute(ctx, string(input))
|
result, err := memstore.NewMemoryTool(t.store, t.agentID, t.currentSession(ctx)).Execute(ctx, string(input))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return tools.ErrorResult("memory tool error: " + err.Error())
|
return tools.ErrorResult("memory tool error: " + err.Error())
|
||||||
}
|
}
|
||||||
|
|
@ -118,11 +118,9 @@ func (t *MemGPTTool) SetSessionResolver(sessionKeyFn func() string) {
|
||||||
t.sessionKeyFn = sessionKeyFn
|
t.sessionKeyFn = sessionKeyFn
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *MemGPTTool) currentSession() string {
|
func (t *MemGPTTool) currentSession(ctx context.Context) string {
|
||||||
if t.sessionKeyFn != nil {
|
if sessionKey := tools.ResolveSessionKey(ctx, t.sessionKeyFn); sessionKey != "" {
|
||||||
if sessionKey := strings.TrimSpace(t.sessionKeyFn()); sessionKey != "" {
|
return sessionKey
|
||||||
return sessionKey
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
if strings.TrimSpace(t.session) != "" {
|
if strings.TrimSpace(t.session) != "" {
|
||||||
return t.session
|
return t.session
|
||||||
|
|
|
||||||
|
|
@ -4,10 +4,12 @@ import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
"strings"
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/ZanzyTHEbar/dragonscale/pkg/bus"
|
"github.com/ZanzyTHEbar/dragonscale/pkg/bus"
|
||||||
"github.com/ZanzyTHEbar/dragonscale/pkg/constants"
|
"github.com/ZanzyTHEbar/dragonscale/pkg/constants"
|
||||||
"github.com/ZanzyTHEbar/dragonscale/pkg/logger"
|
"github.com/ZanzyTHEbar/dragonscale/pkg/logger"
|
||||||
|
"github.com/ZanzyTHEbar/dragonscale/pkg/tools"
|
||||||
"github.com/ZanzyTHEbar/dragonscale/pkg/utils"
|
"github.com/ZanzyTHEbar/dragonscale/pkg/utils"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -24,6 +26,9 @@ func (al *AgentLoop) ProcessDirect(ctx context.Context, content, sessionKey stri
|
||||||
}
|
}
|
||||||
|
|
||||||
func (al *AgentLoop) ProcessDirectWithChannel(ctx context.Context, content, sessionKey, channel, chatID string) (string, error) {
|
func (al *AgentLoop) ProcessDirectWithChannel(ctx context.Context, content, sessionKey, channel, chatID string) (string, error) {
|
||||||
|
roundTracker := tools.NewMessageSendTracker()
|
||||||
|
ctx = tools.WithMessageSendTracker(ctx, roundTracker)
|
||||||
|
|
||||||
msg := bus.InboundMessage{
|
msg := bus.InboundMessage{
|
||||||
Channel: channel,
|
Channel: channel,
|
||||||
SenderID: "cron",
|
SenderID: "cron",
|
||||||
|
|
@ -32,12 +37,20 @@ func (al *AgentLoop) ProcessDirectWithChannel(ctx context.Context, content, sess
|
||||||
SessionKey: sessionKey,
|
SessionKey: sessionKey,
|
||||||
}
|
}
|
||||||
|
|
||||||
return al.processMessage(ctx, msg)
|
response, err := al.processMessage(ctx, msg)
|
||||||
|
if err != nil || response == "" || roundTracker.Sent() {
|
||||||
|
return response, err
|
||||||
|
}
|
||||||
|
al.bus.PublishOutbound(bus.OutboundMessage{Channel: channel, ChatID: chatID, Content: response})
|
||||||
|
return response, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// ProcessDirectStreaming processes a message with streaming token delivery.
|
// ProcessDirectStreaming processes a message with streaming token delivery.
|
||||||
// Text deltas are published to the bus as StreamDelta messages in real time.
|
// Text deltas are published to the bus as StreamDelta messages in real time.
|
||||||
func (al *AgentLoop) ProcessDirectStreaming(ctx context.Context, content, sessionKey, channel, chatID string) (string, error) {
|
func (al *AgentLoop) ProcessDirectStreaming(ctx context.Context, content, sessionKey, channel, chatID string) (string, error) {
|
||||||
|
roundTracker := tools.NewMessageSendTracker()
|
||||||
|
ctx = tools.WithMessageSendTracker(ctx, roundTracker)
|
||||||
|
|
||||||
msg := bus.InboundMessage{
|
msg := bus.InboundMessage{
|
||||||
Channel: channel,
|
Channel: channel,
|
||||||
SenderID: "user",
|
SenderID: "user",
|
||||||
|
|
@ -46,6 +59,12 @@ func (al *AgentLoop) ProcessDirectStreaming(ctx context.Context, content, sessio
|
||||||
SessionKey: sessionKey,
|
SessionKey: sessionKey,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if msg.SessionKey != "" && msg.SessionKey != "heartbeat" && !strings.HasPrefix(msg.SessionKey, "heartbeat:") {
|
||||||
|
if err := al.state.SetLastSessionKeyForTarget(ctx, msg.Channel, msg.ChatID, msg.SessionKey); err != nil {
|
||||||
|
logger.WarnCF("agent", "Failed to record last session key", map[string]interface{}{"error": err.Error(), "session_key": msg.SessionKey})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
return al.runAgentLoop(ctx, processOptions{
|
return al.runAgentLoop(ctx, processOptions{
|
||||||
SessionKey: msg.SessionKey,
|
SessionKey: msg.SessionKey,
|
||||||
Channel: msg.Channel,
|
Channel: msg.Channel,
|
||||||
|
|
@ -63,15 +82,18 @@ func (al *AgentLoop) ProcessDirectStreaming(ctx context.Context, content, sessio
|
||||||
// It injects the active session's summary so the agent has awareness of
|
// It injects the active session's summary so the agent has awareness of
|
||||||
// recent user conversation context.
|
// recent user conversation context.
|
||||||
func (al *AgentLoop) ProcessHeartbeat(ctx context.Context, content, channel, chatID string) (string, error) {
|
func (al *AgentLoop) ProcessHeartbeat(ctx context.Context, content, channel, chatID string) (string, error) {
|
||||||
if v := al.activeSessionKey.Load(); v != nil {
|
sourceSessionKey := ""
|
||||||
if key, ok := v.(string); ok && key != "" {
|
if al.state != nil {
|
||||||
if summary := al.sessions.GetSummary(key); summary != "" {
|
sourceSessionKey = al.state.GetLastSessionKeyForTarget(channel, chatID)
|
||||||
content = content + "\n\n## Recent User Context\n" + summary
|
}
|
||||||
}
|
if sourceSessionKey != "" {
|
||||||
|
if summary := al.sessions.GetSummary(sourceSessionKey); summary != "" {
|
||||||
|
content = content + "\n\n## Recent User Context\n" + summary
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
heartbeatSessionKey := fmt.Sprintf("heartbeat:%d", time.Now().UnixNano())
|
||||||
return al.runAgentLoop(ctx, processOptions{
|
return al.runAgentLoop(ctx, processOptions{
|
||||||
SessionKey: "heartbeat",
|
SessionKey: heartbeatSessionKey,
|
||||||
Channel: channel,
|
Channel: channel,
|
||||||
ChatID: chatID,
|
ChatID: chatID,
|
||||||
UserMessage: content,
|
UserMessage: content,
|
||||||
|
|
@ -108,6 +130,11 @@ func (al *AgentLoop) processMessage(ctx context.Context, msg bus.InboundMessage)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Process as user message
|
// Process as user message
|
||||||
|
if msg.SessionKey != "" && msg.SessionKey != "heartbeat" && !strings.HasPrefix(msg.SessionKey, "heartbeat:") {
|
||||||
|
if err := al.state.SetLastSessionKeyForTarget(ctx, msg.Channel, msg.ChatID, msg.SessionKey); err != nil {
|
||||||
|
logger.WarnCF("agent", "Failed to record last session key", map[string]interface{}{"error": err.Error(), "session_key": msg.SessionKey})
|
||||||
|
}
|
||||||
|
}
|
||||||
return al.runAgentLoop(ctx, processOptions{
|
return al.runAgentLoop(ctx, processOptions{
|
||||||
SessionKey: msg.SessionKey,
|
SessionKey: msg.SessionKey,
|
||||||
Channel: msg.Channel,
|
Channel: msg.Channel,
|
||||||
|
|
|
||||||
|
|
@ -35,31 +35,21 @@ type OffloadingToolRuntime struct {
|
||||||
ChunkChars int
|
ChunkChars int
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type persistedToolResult struct {
|
||||||
|
ToolCall fantasy.ToolCallContent
|
||||||
|
Result fantasy.ToolResultContent
|
||||||
|
}
|
||||||
|
|
||||||
func (r OffloadingToolRuntime) Execute(ctx context.Context, tools []fantasy.AgentTool, toolCalls []fantasy.ToolCallContent, _ func(result fantasy.ToolResultContent) error) ([]fantasy.ToolResultContent, error) {
|
func (r OffloadingToolRuntime) Execute(ctx context.Context, tools []fantasy.AgentTool, toolCalls []fantasy.ToolCallContent, _ func(result fantasy.ToolResultContent) error) ([]fantasy.ToolResultContent, error) {
|
||||||
if len(toolCalls) == 0 {
|
if len(toolCalls) == 0 {
|
||||||
return nil, nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
if err := r.validatePersistenceConfig(); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
if r.Base == nil {
|
if r.Base == nil {
|
||||||
r.Base = fantasy.DAGToolRuntime{MaxConcurrency: defaultToolMaxConcurrency}
|
r.Base = fantasy.DAGToolRuntime{MaxConcurrency: defaultToolMaxConcurrency}
|
||||||
}
|
}
|
||||||
if r.KV == nil {
|
|
||||||
return nil, dserrors.New(dserrors.CodeFailedPrecondition, "KV delegate is nil")
|
|
||||||
}
|
|
||||||
if r.Queries == nil {
|
|
||||||
return nil, dserrors.New(dserrors.CodeFailedPrecondition, "db queries is nil")
|
|
||||||
}
|
|
||||||
if r.ConversationID.IsZero() || r.RunID.IsZero() {
|
|
||||||
return nil, dserrors.New(dserrors.CodeInvalidArgument, "conversation_id/run_id is required")
|
|
||||||
}
|
|
||||||
|
|
||||||
threshold := r.ThresholdChars
|
|
||||||
if threshold <= 0 {
|
|
||||||
threshold = 4_000
|
|
||||||
}
|
|
||||||
chunkChars := r.ChunkChars
|
|
||||||
if chunkChars <= 0 {
|
|
||||||
chunkChars = 2_000
|
|
||||||
}
|
|
||||||
|
|
||||||
stepIndex := fantasy.StepIndexFromCtx(ctx)
|
stepIndex := fantasy.StepIndexFromCtx(ctx)
|
||||||
|
|
||||||
|
|
@ -67,9 +57,53 @@ func (r OffloadingToolRuntime) Execute(ctx context.Context, tools []fantasy.Agen
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
if len(results) != len(toolCalls) {
|
||||||
|
return nil, dserrors.New(dserrors.CodeFailedPrecondition, "tool runtime returned mismatched results length")
|
||||||
|
}
|
||||||
|
|
||||||
|
entries := make([]persistedToolResult, len(results))
|
||||||
for i := range results {
|
for i := range results {
|
||||||
tc := toolCalls[i]
|
entries[i] = persistedToolResult{ToolCall: toolCalls[i], Result: results[i]}
|
||||||
|
}
|
||||||
|
|
||||||
|
results, err = r.persistResults(ctx, stepIndex, entries)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return results, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r OffloadingToolRuntime) PersistResults(ctx context.Context, stepIndex int, toolCalls []fantasy.ToolCallContent, results []fantasy.ToolResultContent) ([]fantasy.ToolResultContent, error) {
|
||||||
|
if err := r.validatePersistenceConfig(); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if len(toolCalls) != len(results) {
|
||||||
|
return nil, dserrors.New(dserrors.CodeInvalidArgument, "toolCalls/results length mismatch")
|
||||||
|
}
|
||||||
|
entries := make([]persistedToolResult, len(results))
|
||||||
|
for i := range results {
|
||||||
|
entries[i] = persistedToolResult{ToolCall: toolCalls[i], Result: results[i]}
|
||||||
|
}
|
||||||
|
return r.persistResults(ctx, stepIndex, entries)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r OffloadingToolRuntime) persistResults(ctx context.Context, stepIndex int, entries []persistedToolResult) ([]fantasy.ToolResultContent, error) {
|
||||||
|
if len(entries) == 0 {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
threshold := r.ThresholdChars
|
||||||
|
if threshold <= 0 {
|
||||||
|
threshold = 4_000
|
||||||
|
}
|
||||||
|
chunkChars := r.ChunkChars
|
||||||
|
if chunkChars <= 0 {
|
||||||
|
chunkChars = 2_000
|
||||||
|
}
|
||||||
|
results := make([]fantasy.ToolResultContent, len(entries))
|
||||||
|
for i, entry := range entries {
|
||||||
|
results[i] = entry.Result
|
||||||
|
tc := entry.ToolCall
|
||||||
res := results[i]
|
res := results[i]
|
||||||
|
|
||||||
fullKey := toolResultFullKey(r.ConversationID, r.RunID, stepIndex, tc.ToolCallID)
|
fullKey := toolResultFullKey(r.ConversationID, r.RunID, stepIndex, tc.ToolCallID)
|
||||||
|
|
@ -147,6 +181,19 @@ func (r OffloadingToolRuntime) Execute(ctx context.Context, tools []fantasy.Agen
|
||||||
return results, nil
|
return results, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (r OffloadingToolRuntime) validatePersistenceConfig() error {
|
||||||
|
if r.KV == nil {
|
||||||
|
return dserrors.New(dserrors.CodeFailedPrecondition, "KV delegate is nil")
|
||||||
|
}
|
||||||
|
if r.Queries == nil {
|
||||||
|
return dserrors.New(dserrors.CodeFailedPrecondition, "db queries is nil")
|
||||||
|
}
|
||||||
|
if r.ConversationID.IsZero() || r.RunID.IsZero() {
|
||||||
|
return dserrors.New(dserrors.CodeInvalidArgument, "conversation_id/run_id is required")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func toolResultFullKey(conversationID, runID ids.UUID, stepIndex int, toolCallID string) string {
|
func toolResultFullKey(conversationID, runID ids.UUID, stepIndex int, toolCallID string) string {
|
||||||
return "tool_results/" + conversationID.String() + "/" + runID.String() + "/step_" + strconv.Itoa(stepIndex) + "/" + sanitizeKeyPart(toolCallID) + "/full.json"
|
return "tool_results/" + conversationID.String() + "/" + runID.String() + "/step_" + strconv.Itoa(stepIndex) + "/" + sanitizeKeyPart(toolCallID) + "/full.json"
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -12,6 +12,7 @@ import (
|
||||||
"github.com/ZanzyTHEbar/dragonscale/pkg/bus"
|
"github.com/ZanzyTHEbar/dragonscale/pkg/bus"
|
||||||
"github.com/ZanzyTHEbar/dragonscale/pkg/config"
|
"github.com/ZanzyTHEbar/dragonscale/pkg/config"
|
||||||
memsqlc "github.com/ZanzyTHEbar/dragonscale/pkg/memory/sqlc"
|
memsqlc "github.com/ZanzyTHEbar/dragonscale/pkg/memory/sqlc"
|
||||||
|
"github.com/ZanzyTHEbar/dragonscale/pkg/tools"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
)
|
)
|
||||||
|
|
@ -159,6 +160,28 @@ func (m *sameStepMultiToolModel) Provider() string { return "mock" }
|
||||||
|
|
||||||
func (m *sameStepMultiToolModel) Model() string { return "same-step-multi-tool-model" }
|
func (m *sameStepMultiToolModel) Model() string { return "same-step-multi-tool-model" }
|
||||||
|
|
||||||
|
type failingModel struct{}
|
||||||
|
|
||||||
|
func (m *failingModel) Generate(_ context.Context, _ fantasy.Call) (*fantasy.Response, error) {
|
||||||
|
return nil, fmt.Errorf("synthetic generate failure")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *failingModel) Stream(_ context.Context, _ fantasy.Call) (fantasy.StreamResponse, error) {
|
||||||
|
return nil, fmt.Errorf("synthetic stream failure")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *failingModel) GenerateObject(_ context.Context, _ fantasy.ObjectCall) (*fantasy.ObjectResponse, error) {
|
||||||
|
return nil, fmt.Errorf("not implemented")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *failingModel) StreamObject(_ context.Context, _ fantasy.ObjectCall) (fantasy.ObjectStreamResponse, error) {
|
||||||
|
return nil, fmt.Errorf("not implemented")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *failingModel) Provider() string { return "mock" }
|
||||||
|
|
||||||
|
func (m *failingModel) Model() string { return "failing-model" }
|
||||||
|
|
||||||
func countPromptToolResults(prompt []fantasy.Message) int {
|
func countPromptToolResults(prompt []fantasy.Message) int {
|
||||||
count := 0
|
count := 0
|
||||||
for _, msg := range prompt {
|
for _, msg := range prompt {
|
||||||
|
|
@ -184,6 +207,19 @@ func uniqueTransitionSteps(rows []memsqlc.AgentStateTransition) []int64 {
|
||||||
return steps
|
return steps
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func uniqueRunStateSteps(rows []memsqlc.AgentRunState) []int64 {
|
||||||
|
seen := make(map[int64]struct{})
|
||||||
|
steps := make([]int64, 0, len(rows))
|
||||||
|
for _, row := range rows {
|
||||||
|
if _, ok := seen[row.StepIndex]; ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
seen[row.StepIndex] = struct{}{}
|
||||||
|
steps = append(steps, row.StepIndex)
|
||||||
|
}
|
||||||
|
return steps
|
||||||
|
}
|
||||||
|
|
||||||
func TestIntegration_RuntimeBookkeeping_PersistsTransitionsAndMetrics(t *testing.T) {
|
func TestIntegration_RuntimeBookkeeping_PersistsTransitionsAndMetrics(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
|
|
@ -249,6 +285,14 @@ func TestIntegration_RuntimeBookkeeping_PersistsTransitionsAndMetrics(t *testing
|
||||||
require.Len(t, toolResults, 2)
|
require.Len(t, toolResults, 2)
|
||||||
assert.Equal(t, int64(0), toolResults[0].StepIndex)
|
assert.Equal(t, int64(0), toolResults[0].StepIndex)
|
||||||
assert.Equal(t, int64(1), toolResults[1].StepIndex)
|
assert.Equal(t, int64(1), toolResults[1].StepIndex)
|
||||||
|
|
||||||
|
runStates, err := al.queries.ListAgentRunStatesByRunID(ctx, memsqlc.ListAgentRunStatesByRunIDParams{
|
||||||
|
RunID: completion.RunID,
|
||||||
|
Lim: 16,
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, runStates, 3)
|
||||||
|
assert.Equal(t, []int64{0, 1, 3}, uniqueRunStateSteps(runStates))
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestIntegration_RuntimeBookkeeping_UsesAgentStepForMultipleToolCalls(t *testing.T) {
|
func TestIntegration_RuntimeBookkeeping_UsesAgentStepForMultipleToolCalls(t *testing.T) {
|
||||||
|
|
@ -308,4 +352,177 @@ func TestIntegration_RuntimeBookkeeping_UsesAgentStepForMultipleToolCalls(t *tes
|
||||||
require.Len(t, toolResults, 2)
|
require.Len(t, toolResults, 2)
|
||||||
assert.Equal(t, int64(0), toolResults[0].StepIndex)
|
assert.Equal(t, int64(0), toolResults[0].StepIndex)
|
||||||
assert.Equal(t, int64(0), toolResults[1].StepIndex)
|
assert.Equal(t, int64(0), toolResults[1].StepIndex)
|
||||||
|
|
||||||
|
runStates, err := al.queries.ListAgentRunStatesByRunID(ctx, memsqlc.ListAgentRunStatesByRunIDParams{
|
||||||
|
RunID: completion.RunID,
|
||||||
|
Lim: 16,
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, runStates, 2)
|
||||||
|
assert.Equal(t, []int64{0, 2}, uniqueRunStateSteps(runStates))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIntegration_RuntimeBookkeeping_FailedRunIsTerminalized(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
tmpDir, err := os.MkdirTemp("", "agent-runtime-failed-run-*")
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer os.RemoveAll(tmpDir)
|
||||||
|
|
||||||
|
cfg := &config.Config{
|
||||||
|
Agents: config.AgentsConfig{
|
||||||
|
Defaults: config.AgentDefaults{
|
||||||
|
Sandbox: tmpDir,
|
||||||
|
Model: "failing-model",
|
||||||
|
MaxTokens: 4096,
|
||||||
|
MaxToolIterations: 10,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
msgBus := bus.NewMessageBus()
|
||||||
|
al := mustNewAgentLoop(t, cfg, msgBus, &failingModel{})
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(t.Context(), 10*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
msg := bus.InboundMessage{
|
||||||
|
Channel: "test",
|
||||||
|
SenderID: "user1",
|
||||||
|
ChatID: "chat1",
|
||||||
|
Content: "trigger failure",
|
||||||
|
SessionKey: "runtime-bookkeeping-failed",
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err = al.processMessage(ctx, msg)
|
||||||
|
require.Error(t, err)
|
||||||
|
|
||||||
|
convID, ok := al.conversationIDs.Load(msg.SessionKey)
|
||||||
|
require.True(t, ok)
|
||||||
|
run, err := al.queries.GetLatestAgentRunByConversationID(ctx, memsqlc.GetLatestAgentRunByConversationIDParams{
|
||||||
|
ConversationID: convID,
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, "failed", run.Status)
|
||||||
|
assert.Contains(t, string(run.MetadataJson), "synthetic generate failure")
|
||||||
|
assert.Contains(t, string(run.MetadataJson), `"reason":"failed"`)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIntegration_RuntimeBookkeeping_SubagentRunIsTerminalized(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
tmpDir, err := os.MkdirTemp("", "agent-runtime-subagent-run-*")
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer os.RemoveAll(tmpDir)
|
||||||
|
|
||||||
|
cfg := &config.Config{
|
||||||
|
Agents: config.AgentsConfig{
|
||||||
|
Defaults: config.AgentDefaults{
|
||||||
|
Sandbox: tmpDir,
|
||||||
|
Model: "mock-model",
|
||||||
|
MaxTokens: 4096,
|
||||||
|
MaxToolIterations: 10,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
msgBus := bus.NewMessageBus()
|
||||||
|
al := mustNewAgentLoop(t, cfg, msgBus, newMockLanguageModel("subagent final response"))
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(t.Context(), 10*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
runLoop := MakeUnifiedRunLoopFunc(al)
|
||||||
|
result, err := runLoop(ctx, tools.ToolLoopConfig{Model: newMockLanguageModel("subagent final response"), MaxIterations: 3}, "", "subagent task", "test", "chat1")
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, result)
|
||||||
|
assert.Equal(t, "subagent final response", result.Content)
|
||||||
|
|
||||||
|
conversations, err := al.queries.ListAgentConversations(ctx, memsqlc.ListAgentConversationsParams{Limit: 10})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, conversations, 1)
|
||||||
|
|
||||||
|
latestRun, err := al.queries.GetLatestAgentRunByConversationID(ctx, memsqlc.GetLatestAgentRunByConversationIDParams{
|
||||||
|
ConversationID: conversations[0].ID,
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.False(t, latestRun.ID.IsZero())
|
||||||
|
assert.Equal(t, "completed", latestRun.Status)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIntegration_RuntimeBookkeeping_SubagentFailedRunIsTerminalized(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
tmpDir, err := os.MkdirTemp("", "agent-runtime-subagent-failed-*")
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer os.RemoveAll(tmpDir)
|
||||||
|
|
||||||
|
cfg := &config.Config{
|
||||||
|
Agents: config.AgentsConfig{
|
||||||
|
Defaults: config.AgentDefaults{
|
||||||
|
Sandbox: tmpDir,
|
||||||
|
Model: "failing-model",
|
||||||
|
MaxTokens: 4096,
|
||||||
|
MaxToolIterations: 10,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
msgBus := bus.NewMessageBus()
|
||||||
|
al := mustNewAgentLoop(t, cfg, msgBus, &failingModel{})
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(t.Context(), 10*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
runLoop := MakeUnifiedRunLoopFunc(al)
|
||||||
|
_, err = runLoop(ctx, tools.ToolLoopConfig{Model: &failingModel{}, MaxIterations: 3}, "", "subagent failing task", "test", "chat1")
|
||||||
|
require.Error(t, err)
|
||||||
|
|
||||||
|
conversations, err := al.queries.ListAgentConversations(ctx, memsqlc.ListAgentConversationsParams{Limit: 10})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, conversations, 1)
|
||||||
|
|
||||||
|
latestRun, err := al.queries.GetLatestAgentRunByConversationID(ctx, memsqlc.GetLatestAgentRunByConversationIDParams{
|
||||||
|
ConversationID: conversations[0].ID,
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.False(t, latestRun.ID.IsZero())
|
||||||
|
assert.Equal(t, "failed", latestRun.Status)
|
||||||
|
assert.Contains(t, string(latestRun.MetadataJson), "synthetic generate failure")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIntegration_RuntimeBookkeeping_SubagentUsesDelegatedParentSession(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
tmpDir, err := os.MkdirTemp("", "agent-runtime-subagent-session-*")
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer os.RemoveAll(tmpDir)
|
||||||
|
|
||||||
|
cfg := &config.Config{
|
||||||
|
Agents: config.AgentsConfig{
|
||||||
|
Defaults: config.AgentDefaults{
|
||||||
|
Sandbox: tmpDir,
|
||||||
|
Model: "mock-model",
|
||||||
|
MaxTokens: 4096,
|
||||||
|
MaxToolIterations: 10,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
msgBus := bus.NewMessageBus()
|
||||||
|
al := mustNewAgentLoop(t, cfg, msgBus, newMockLanguageModel("subagent final response"))
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(t.Context(), 10*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
runLoop := MakeUnifiedRunLoopFunc(al)
|
||||||
|
parentCtx := tools.WithSessionKey(ctx, "parent-session")
|
||||||
|
_, err = runLoop(parentCtx, tools.ToolLoopConfig{Model: newMockLanguageModel("subagent final response"), MaxIterations: 3}, "", "subagent task", "test", "chat1")
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
conversations, err := al.queries.ListAgentConversations(ctx, memsqlc.ListAgentConversationsParams{Limit: 10})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, conversations, 1)
|
||||||
|
require.NotNil(t, conversations[0].Title)
|
||||||
|
assert.Contains(t, *conversations[0].Title, "parent-session::subagent::")
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -12,20 +12,25 @@ import (
|
||||||
"github.com/ZanzyTHEbar/dragonscale/pkg/itr"
|
"github.com/ZanzyTHEbar/dragonscale/pkg/itr"
|
||||||
"github.com/ZanzyTHEbar/dragonscale/pkg/logger"
|
"github.com/ZanzyTHEbar/dragonscale/pkg/logger"
|
||||||
"github.com/ZanzyTHEbar/dragonscale/pkg/security/securebus"
|
"github.com/ZanzyTHEbar/dragonscale/pkg/security/securebus"
|
||||||
|
"github.com/ZanzyTHEbar/dragonscale/pkg/tools"
|
||||||
)
|
)
|
||||||
|
|
||||||
// SecureBusToolRuntime is a fantasy.ToolRuntime that routes every tool call
|
// SecureBusToolRuntime is a fantasy.ToolRuntime that routes every tool call
|
||||||
// through the SecureBus before (and after) passing it to the underlying runtime.
|
// through the SecureBus, then persists the final post-policy result through the
|
||||||
|
// offloading runtime.
|
||||||
//
|
//
|
||||||
// Pipeline per tool call:
|
// Pipeline per tool call:
|
||||||
// 1. Serialize tool call args → ToolRequest
|
// 1. Serialize tool call args → ToolRequest
|
||||||
// 2. bus.Execute() → capability check, secret injection, output scan, audit
|
// 2. bus.Execute() → capability check, secret injection, output scan, audit
|
||||||
// 3. If bus returns a policy error, short-circuit with that error result
|
// 3. If bus returns a policy error, short-circuit with that error result
|
||||||
// 4. Otherwise delegate to Base runtime for actual execution
|
// 4. Persist the final result through the offloading runtime
|
||||||
// 5. If bus detected a leak, replace Base output with the redacted version
|
|
||||||
type SecureBusToolRuntime struct {
|
type SecureBusToolRuntime struct {
|
||||||
// Base is the underlying runtime and is required.
|
// Offloader persists tool results and applies large-result truncation.
|
||||||
Base fantasy.ToolRuntime
|
Offloader OffloadingToolRuntime
|
||||||
|
|
||||||
|
// FantasyTools allows execution of agent-only fantasy tools that are not part
|
||||||
|
// of the raw registry/securebus path.
|
||||||
|
FantasyTools map[string]fantasy.AgentTool
|
||||||
|
|
||||||
// Bus is required.
|
// Bus is required.
|
||||||
Bus *securebus.Bus
|
Bus *securebus.Bus
|
||||||
|
|
@ -33,6 +38,11 @@ type SecureBusToolRuntime struct {
|
||||||
// SessionKey is forwarded to bus requests for audit tracing.
|
// SessionKey is forwarded to bus requests for audit tracing.
|
||||||
SessionKey string
|
SessionKey string
|
||||||
|
|
||||||
|
// Channel and ChatID restore the registry execution context that contextual
|
||||||
|
// and async tools expect on the main runtime path.
|
||||||
|
Channel string
|
||||||
|
ChatID string
|
||||||
|
|
||||||
// UserPrompt allows lightweight repair of placeholder tool arguments when
|
// UserPrompt allows lightweight repair of placeholder tool arguments when
|
||||||
// the user provided an explicit literal value in the request.
|
// the user provided an explicit literal value in the request.
|
||||||
UserPrompt string
|
UserPrompt string
|
||||||
|
|
@ -45,7 +55,7 @@ type SecureBusToolRuntime struct {
|
||||||
// Execute implements fantasy.ToolRuntime.
|
// Execute implements fantasy.ToolRuntime.
|
||||||
func (r SecureBusToolRuntime) Execute(
|
func (r SecureBusToolRuntime) Execute(
|
||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
tools []fantasy.AgentTool,
|
_ []fantasy.AgentTool,
|
||||||
toolCalls []fantasy.ToolCallContent,
|
toolCalls []fantasy.ToolCallContent,
|
||||||
onResult func(fantasy.ToolResultContent) error,
|
onResult func(fantasy.ToolResultContent) error,
|
||||||
) ([]fantasy.ToolResultContent, error) {
|
) ([]fantasy.ToolResultContent, error) {
|
||||||
|
|
@ -55,44 +65,74 @@ func (r SecureBusToolRuntime) Execute(
|
||||||
if r.Bus == nil {
|
if r.Bus == nil {
|
||||||
return nil, fmt.Errorf("secure bus runtime requires bus")
|
return nil, fmt.Errorf("secure bus runtime requires bus")
|
||||||
}
|
}
|
||||||
if r.Base == nil {
|
if r.Offloader.KV == nil {
|
||||||
return nil, fmt.Errorf("secure bus runtime requires base runtime")
|
return nil, fmt.Errorf("secure bus runtime requires offloader")
|
||||||
}
|
}
|
||||||
|
|
||||||
results := make([]fantasy.ToolResultContent, 0, len(toolCalls))
|
results := make([]fantasy.ToolResultContent, 0, len(toolCalls))
|
||||||
stepIndex := fantasy.StepIndexFromCtx(ctx)
|
stepIndex := fantasy.StepIndexFromCtx(ctx)
|
||||||
|
|
||||||
type deferredState struct {
|
var finalState *runtimeStepState
|
||||||
step int
|
defer func() {
|
||||||
state string
|
if finalState != nil {
|
||||||
snapshot map[string]any
|
r.recordRunState(context.WithoutCancel(ctx), finalState.step, finalState.state, finalState.snapshot)
|
||||||
}
|
}
|
||||||
var pendingStates []deferredState
|
}()
|
||||||
|
|
||||||
for i, tc := range toolCalls {
|
for i, tc := range toolCalls {
|
||||||
tc = repairToolCallInput(tc, r.UserPrompt)
|
tc = repairToolCallInput(tc, r.UserPrompt)
|
||||||
step := stepIndex
|
step := stepIndex
|
||||||
pendingStates = append(pendingStates, deferredState{step, "tool_call", map[string]any{
|
execCtx := tools.WithExecutionTarget(ctx, r.Channel, r.ChatID)
|
||||||
|
finalState = &runtimeStepState{step: step, state: "tool_call", snapshot: map[string]any{
|
||||||
"tool_name": tc.ToolName,
|
"tool_name": tc.ToolName,
|
||||||
"tool_call_index": i,
|
"tool_call_index": i,
|
||||||
}})
|
}}
|
||||||
|
|
||||||
|
if ft, ok := r.FantasyTools[tc.ToolName]; ok {
|
||||||
|
tr, err := executeFantasyTool(ctx, ft, tc)
|
||||||
|
if err != nil {
|
||||||
|
return results, err
|
||||||
|
}
|
||||||
|
persisted, err := r.Offloader.PersistResults(fantasy.WithStepIndex(ctx, step), step, []fantasy.ToolCallContent{tc}, []fantasy.ToolResultContent{tr})
|
||||||
|
if err != nil {
|
||||||
|
return results, err
|
||||||
|
}
|
||||||
|
for _, pr := range persisted {
|
||||||
|
results = append(results, pr)
|
||||||
|
finalState = &runtimeStepState{step: step, state: "tool_result", snapshot: map[string]any{
|
||||||
|
"tool_name": tc.ToolName,
|
||||||
|
"tool_call_index": i,
|
||||||
|
}}
|
||||||
|
if onResult != nil {
|
||||||
|
if err := onResult(pr); err != nil {
|
||||||
|
return results, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
reqID := ids.New().String()
|
reqID := ids.New().String()
|
||||||
req := itr.NewToolExecRequest(reqID, r.SessionKey, tc.ToolCallID, tc.ToolName, tc.Input)
|
req := itr.NewToolExecRequest(reqID, r.SessionKey, tc.ToolCallID, tc.ToolName, tc.Input)
|
||||||
busResp := r.Bus.Execute(ctx, req)
|
busResp := r.Bus.Execute(execCtx, req)
|
||||||
|
|
||||||
if busResp.IsError {
|
if busResp.IsError {
|
||||||
sanitized := sanitizePolicyError(busResp.Result)
|
sanitized := sanitizePolicyError(busResp.Result)
|
||||||
tr := fantasy.ToolResultContent{
|
result := fantasy.ToolResultContent{
|
||||||
ToolCallID: tc.ToolCallID,
|
ToolCallID: tc.ToolCallID,
|
||||||
ToolName: tc.ToolName,
|
ToolName: tc.ToolName,
|
||||||
Result: fantasy.ToolResultOutputContentError{Error: errors.New(sanitized)},
|
Result: fantasy.ToolResultOutputContentError{Error: errors.New(sanitized)},
|
||||||
}
|
}
|
||||||
pendingStates = append(pendingStates, deferredState{step, "tool_call_error", map[string]any{
|
persisted, err := r.Offloader.PersistResults(fantasy.WithStepIndex(ctx, step), step, []fantasy.ToolCallContent{tc}, []fantasy.ToolResultContent{result})
|
||||||
|
if err != nil {
|
||||||
|
return results, err
|
||||||
|
}
|
||||||
|
tr := persisted[0]
|
||||||
|
finalState = &runtimeStepState{step: step, state: "tool_call_error", snapshot: map[string]any{
|
||||||
"tool_name": tc.ToolName,
|
"tool_name": tc.ToolName,
|
||||||
"error": busResp.Result,
|
"error": busResp.Result,
|
||||||
"error_safe": sanitized,
|
"error_safe": sanitized,
|
||||||
}})
|
}}
|
||||||
results = append(results, tr)
|
results = append(results, tr)
|
||||||
if onResult != nil {
|
if onResult != nil {
|
||||||
if err := onResult(tr); err != nil {
|
if err := onResult(tr); err != nil {
|
||||||
|
|
@ -102,21 +142,22 @@ func (r SecureBusToolRuntime) Execute(
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
// Execute via Base runtime for the single tool call.
|
result := fantasy.ToolResultContent{
|
||||||
baseResults, err := r.Base.Execute(fantasy.WithStepIndex(ctx, stepIndex), tools, []fantasy.ToolCallContent{tc}, nil)
|
ToolCallID: tc.ToolCallID,
|
||||||
|
ToolName: tc.ToolName,
|
||||||
|
Result: fantasy.ToolResultOutputContentText{Text: busResp.Result},
|
||||||
|
}
|
||||||
|
persisted, err := r.Offloader.PersistResults(fantasy.WithStepIndex(ctx, step), step, []fantasy.ToolCallContent{tc}, []fantasy.ToolResultContent{result})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return results, err
|
return results, err
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, br := range baseResults {
|
for _, br := range persisted {
|
||||||
if busResp.LeakDetected {
|
|
||||||
br = overrideResultText(br, busResp.Result)
|
|
||||||
}
|
|
||||||
results = append(results, br)
|
results = append(results, br)
|
||||||
pendingStates = append(pendingStates, deferredState{step, "tool_result", map[string]any{
|
finalState = &runtimeStepState{step: step, state: "tool_result", snapshot: map[string]any{
|
||||||
"tool_name": tc.ToolName,
|
"tool_name": tc.ToolName,
|
||||||
"tool_call_index": i,
|
"tool_call_index": i,
|
||||||
}})
|
}}
|
||||||
if onResult != nil {
|
if onResult != nil {
|
||||||
if err := onResult(br); err != nil {
|
if err := onResult(br); err != nil {
|
||||||
return results, err
|
return results, err
|
||||||
|
|
@ -125,14 +166,15 @@ func (r SecureBusToolRuntime) Execute(
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Flush all buffered state writes in one pass
|
|
||||||
for _, ps := range pendingStates {
|
|
||||||
r.recordRunState(ctx, ps.step, ps.state, ps.snapshot)
|
|
||||||
}
|
|
||||||
|
|
||||||
return results, nil
|
return results, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type runtimeStepState struct {
|
||||||
|
step int
|
||||||
|
state string
|
||||||
|
snapshot map[string]any
|
||||||
|
}
|
||||||
|
|
||||||
func (r SecureBusToolRuntime) recordRunState(ctx context.Context, stepIndex int, state string, snapshot map[string]any) {
|
func (r SecureBusToolRuntime) recordRunState(ctx context.Context, stepIndex int, state string, snapshot map[string]any) {
|
||||||
if r.StateStore == nil || r.RunID.IsZero() {
|
if r.StateStore == nil || r.RunID.IsZero() {
|
||||||
return
|
return
|
||||||
|
|
@ -182,15 +224,6 @@ func sanitizePolicyError(raw string) string {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// overrideResultText replaces the text output of a ToolResultContent with
|
|
||||||
// the redacted version produced by the SecureBus.
|
|
||||||
func overrideResultText(tr fantasy.ToolResultContent, text string) fantasy.ToolResultContent {
|
|
||||||
if _, ok := tr.Result.(fantasy.ToolResultOutputContentText); ok {
|
|
||||||
tr.Result = fantasy.ToolResultOutputContentText{Text: text}
|
|
||||||
}
|
|
||||||
return tr
|
|
||||||
}
|
|
||||||
|
|
||||||
func repairToolCallInput(tc fantasy.ToolCallContent, userPrompt string) fantasy.ToolCallContent {
|
func repairToolCallInput(tc fantasy.ToolCallContent, userPrompt string) fantasy.ToolCallContent {
|
||||||
if strings.TrimSpace(userPrompt) == "" || strings.TrimSpace(tc.Input) == "" {
|
if strings.TrimSpace(userPrompt) == "" || strings.TrimSpace(tc.Input) == "" {
|
||||||
return tc
|
return tc
|
||||||
|
|
@ -246,6 +279,66 @@ func repairToolCallInput(tc fantasy.ToolCallContent, userPrompt string) fantasy.
|
||||||
return tc
|
return tc
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func executeFantasyTool(ctx context.Context, tool fantasy.AgentTool, toolCall fantasy.ToolCallContent) (fantasy.ToolResultContent, error) {
|
||||||
|
result := fantasy.ToolResultContent{
|
||||||
|
ToolCallID: toolCall.ToolCallID,
|
||||||
|
ToolName: toolCall.ToolName,
|
||||||
|
ProviderExecuted: false,
|
||||||
|
}
|
||||||
|
|
||||||
|
response, err := tool.Run(ctx, fantasy.ToolCall{
|
||||||
|
ID: toolCall.ToolCallID,
|
||||||
|
Name: toolCall.ToolName,
|
||||||
|
Input: toolCall.Input,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return result, err
|
||||||
|
}
|
||||||
|
|
||||||
|
result.ClientMetadata = response.Metadata
|
||||||
|
if response.IsError {
|
||||||
|
result.Result = fantasy.ToolResultOutputContentError{Error: errors.New(response.Content)}
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
switch response.Type {
|
||||||
|
case "image", "media":
|
||||||
|
result.Result = fantasy.ToolResultOutputContentMedia{
|
||||||
|
Data: string(response.Data),
|
||||||
|
MediaType: response.MediaType,
|
||||||
|
Text: response.Content,
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
result.Result = fantasy.ToolResultOutputContentText{Text: response.Content}
|
||||||
|
}
|
||||||
|
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func fantasyToolMap(toolsList []fantasy.AgentTool, registry *tools.ToolRegistry) map[string]fantasy.AgentTool {
|
||||||
|
if len(toolsList) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
fallback := make(map[string]fantasy.AgentTool)
|
||||||
|
for _, tool := range toolsList {
|
||||||
|
if tool == nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
name := tool.Info().Name
|
||||||
|
if registry != nil {
|
||||||
|
if _, ok := registry.Get(name); ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
}
|
||||||
|
fallback[name] = tool
|
||||||
|
}
|
||||||
|
if len(fallback) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return fallback
|
||||||
|
}
|
||||||
|
|
||||||
func repairDirectArg(input, field, replacement string) (string, bool) {
|
func repairDirectArg(input, field, replacement string) (string, bool) {
|
||||||
var args map[string]any
|
var args map[string]any
|
||||||
if err := json.Unmarshal([]byte(input), &args); err != nil {
|
if err := json.Unmarshal([]byte(input), &args); err != nil {
|
||||||
|
|
|
||||||
|
|
@ -1,12 +1,107 @@
|
||||||
package agent
|
package agent
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"strings"
|
"strings"
|
||||||
|
"sync/atomic"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
fantasy "charm.land/fantasy"
|
fantasy "charm.land/fantasy"
|
||||||
|
"github.com/ZanzyTHEbar/dragonscale/pkg/ids"
|
||||||
|
"github.com/ZanzyTHEbar/dragonscale/pkg/memory/delegate"
|
||||||
|
"github.com/ZanzyTHEbar/dragonscale/pkg/memory/sqlc"
|
||||||
|
"github.com/ZanzyTHEbar/dragonscale/pkg/security/securebus"
|
||||||
|
"github.com/ZanzyTHEbar/dragonscale/pkg/tools"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
type secureBusTestDB struct {
|
||||||
|
delegate *delegate.LibSQLDelegate
|
||||||
|
}
|
||||||
|
|
||||||
|
func newSecureBusTestDB(t *testing.T) *secureBusTestDB {
|
||||||
|
t.Helper()
|
||||||
|
d, err := delegate.NewLibSQLInMemory()
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, d.Init(t.Context()))
|
||||||
|
t.Cleanup(func() { _ = d.Close() })
|
||||||
|
return &secureBusTestDB{delegate: d}
|
||||||
|
}
|
||||||
|
|
||||||
|
func newSecureBusConversation(t *testing.T, q *sqlc.Queries) ids.UUID {
|
||||||
|
t.Helper()
|
||||||
|
id := ids.New()
|
||||||
|
title := "securebus-test-conv"
|
||||||
|
_, err := q.CreateAgentConversation(t.Context(), sqlc.CreateAgentConversationParams{
|
||||||
|
ID: id,
|
||||||
|
Title: &title,
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
return id
|
||||||
|
}
|
||||||
|
|
||||||
|
type countingTool struct {
|
||||||
|
calls atomic.Int32
|
||||||
|
text string
|
||||||
|
err bool
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *countingTool) Name() string { return "echo" }
|
||||||
|
func (t *countingTool) Description() string { return "echo" }
|
||||||
|
func (t *countingTool) Parameters() map[string]interface{} {
|
||||||
|
return map[string]interface{}{"type": "object", "properties": map[string]interface{}{}}
|
||||||
|
}
|
||||||
|
func (t *countingTool) Execute(_ context.Context, args map[string]interface{}) *tools.ToolResult {
|
||||||
|
t.calls.Add(1)
|
||||||
|
text, _ := args["text"].(string)
|
||||||
|
if t.text != "" {
|
||||||
|
text = t.text
|
||||||
|
}
|
||||||
|
return &tools.ToolResult{ForLLM: text, IsError: t.err}
|
||||||
|
}
|
||||||
|
|
||||||
|
func makeSecureBusRuntimeFixture(t *testing.T, tool tools.Tool, policy securebus.PolicyConfig) (SecureBusToolRuntime, *sqlc.Queries, KVDelegate, ids.UUID) {
|
||||||
|
t.Helper()
|
||||||
|
db := newSecureBusTestDB(t)
|
||||||
|
q := db.delegate.Queries()
|
||||||
|
convID := newSecureBusConversation(t, q)
|
||||||
|
stateStore := NewStateStore(q)
|
||||||
|
run, err := stateStore.CreateRun(t.Context(), convID)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
kv := NewDelegateKV(db.delegate, "securebus-runtime-test")
|
||||||
|
capLookup := func(name string) (tools.ToolCapabilities, bool) {
|
||||||
|
if name != tool.Name() {
|
||||||
|
return tools.ZeroCapabilities(), false
|
||||||
|
}
|
||||||
|
return tools.ExtractCapabilities(tool), true
|
||||||
|
}
|
||||||
|
executor := func(ctx context.Context, name string, args map[string]interface{}) *tools.ToolResult {
|
||||||
|
if name != tool.Name() {
|
||||||
|
return &tools.ToolResult{ForLLM: "tool not found", IsError: true}
|
||||||
|
}
|
||||||
|
return tool.Execute(ctx, args)
|
||||||
|
}
|
||||||
|
bus := securebus.New(securebus.BusConfig{Policy: policy, Workers: 1}, nil, capLookup, executor)
|
||||||
|
t.Cleanup(bus.Close)
|
||||||
|
|
||||||
|
return SecureBusToolRuntime{
|
||||||
|
Offloader: OffloadingToolRuntime{
|
||||||
|
KV: kv,
|
||||||
|
Queries: q,
|
||||||
|
ConversationID: convID,
|
||||||
|
RunID: run.ID,
|
||||||
|
ThresholdChars: 4_000,
|
||||||
|
ChunkChars: 2_000,
|
||||||
|
},
|
||||||
|
Bus: bus,
|
||||||
|
SessionKey: "securebus-test-session",
|
||||||
|
StateStore: stateStore,
|
||||||
|
RunID: run.ID,
|
||||||
|
}, q, kv, run.ID
|
||||||
|
}
|
||||||
|
|
||||||
func TestRepairToolCallInputRepairsDirectExecPlaceholder(t *testing.T) {
|
func TestRepairToolCallInputRepairsDirectExecPlaceholder(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
|
|
@ -131,3 +226,139 @@ func TestSanitizePolicyErrorRedactsPolicyViolations(t *testing.T) {
|
||||||
t.Fatalf("expected redacted policy text, got %q", got)
|
t.Fatalf("expected redacted policy text, got %q", got)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestSecureBusToolRuntime_ExecutesToolExactlyOnce(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
tool := &countingTool{}
|
||||||
|
runtime, _, _, _ := makeSecureBusRuntimeFixture(t, tool, securebus.DefaultPolicyConfig())
|
||||||
|
|
||||||
|
results, err := runtime.Execute(
|
||||||
|
fantasy.WithStepIndex(t.Context(), 0),
|
||||||
|
nil,
|
||||||
|
[]fantasy.ToolCallContent{{ToolCallID: "call-1", ToolName: "echo", Input: `{"text":"hello"}`}},
|
||||||
|
nil,
|
||||||
|
)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, results, 1)
|
||||||
|
assert.Equal(t, int32(1), tool.calls.Load())
|
||||||
|
assert.Equal(t, "hello", results[0].Result.(fantasy.ToolResultOutputContentText).Text)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSecureBusToolRuntime_ExecutesEachToolCallOnce(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
tool := &countingTool{}
|
||||||
|
runtime, _, _, _ := makeSecureBusRuntimeFixture(t, tool, securebus.DefaultPolicyConfig())
|
||||||
|
|
||||||
|
results, err := runtime.Execute(
|
||||||
|
fantasy.WithStepIndex(t.Context(), 0),
|
||||||
|
nil,
|
||||||
|
[]fantasy.ToolCallContent{
|
||||||
|
{ToolCallID: "call-1", ToolName: "echo", Input: `{"text":"one"}`},
|
||||||
|
{ToolCallID: "call-2", ToolName: "echo", Input: `{"text":"two"}`},
|
||||||
|
},
|
||||||
|
nil,
|
||||||
|
)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, results, 2)
|
||||||
|
assert.Equal(t, int32(2), tool.calls.Load())
|
||||||
|
assert.Equal(t, "one", results[0].Result.(fantasy.ToolResultOutputContentText).Text)
|
||||||
|
assert.Equal(t, "two", results[1].Result.(fantasy.ToolResultOutputContentText).Text)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSecureBusToolRuntime_LeakRedactionIsPersisted(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
secret := "AKIAIOSFODNN7EXAMPLE"
|
||||||
|
tool := &countingTool{text: "result: " + secret}
|
||||||
|
runtime, q, kv, runID := makeSecureBusRuntimeFixture(t, tool, securebus.DefaultPolicyConfig())
|
||||||
|
|
||||||
|
results, err := runtime.Execute(
|
||||||
|
fantasy.WithStepIndex(t.Context(), 0),
|
||||||
|
nil,
|
||||||
|
[]fantasy.ToolCallContent{{ToolCallID: "call-1", ToolName: "echo", Input: `{"text":"ignored"}`}},
|
||||||
|
nil,
|
||||||
|
)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, results, 1)
|
||||||
|
assert.Equal(t, int32(1), tool.calls.Load())
|
||||||
|
|
||||||
|
textResult, ok := results[0].Result.(fantasy.ToolResultOutputContentText)
|
||||||
|
require.True(t, ok)
|
||||||
|
assert.NotContains(t, textResult.Text, secret)
|
||||||
|
|
||||||
|
rows, err := q.ListAgentToolResultsByRunID(t.Context(), sqlc.ListAgentToolResultsByRunIDParams{RunID: runID, Lim: 10})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, rows, 1)
|
||||||
|
require.NotNil(t, rows[0].Preview)
|
||||||
|
assert.NotContains(t, *rows[0].Preview, secret)
|
||||||
|
|
||||||
|
persisted, err := kv.Get(t.Context(), rows[0].FullKey)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.NotContains(t, string(persisted), secret)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSecureBusToolRuntime_PersistsToolExecutionErrorResult(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
tool := &countingTool{text: "boom", err: true}
|
||||||
|
runtime, q, kv, runID := makeSecureBusRuntimeFixture(t, tool, securebus.DefaultPolicyConfig())
|
||||||
|
|
||||||
|
results, err := runtime.Execute(
|
||||||
|
fantasy.WithStepIndex(t.Context(), 2),
|
||||||
|
nil,
|
||||||
|
[]fantasy.ToolCallContent{{ToolCallID: "call-1", ToolName: "echo", Input: `{"text":"hello"}`}},
|
||||||
|
nil,
|
||||||
|
)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, results, 1)
|
||||||
|
assert.Equal(t, int32(1), tool.calls.Load())
|
||||||
|
|
||||||
|
errorResult, ok := results[0].Result.(fantasy.ToolResultOutputContentError)
|
||||||
|
require.True(t, ok)
|
||||||
|
assert.Equal(t, "tool execution denied", errorResult.Error.Error())
|
||||||
|
|
||||||
|
rows, err := q.ListAgentToolResultsByRunID(t.Context(), sqlc.ListAgentToolResultsByRunIDParams{RunID: runID, Lim: 10})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, rows, 1)
|
||||||
|
require.NotNil(t, rows[0].Preview)
|
||||||
|
assert.Contains(t, *rows[0].Preview, "tool execution denied")
|
||||||
|
|
||||||
|
persisted, err := kv.Get(t.Context(), rows[0].FullKey)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Contains(t, string(persisted), `"type":"error"`)
|
||||||
|
assert.Contains(t, string(persisted), "tool execution denied")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSecureBusToolRuntime_PersistsInvalidArgsErrorResult(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
tool := &countingTool{}
|
||||||
|
runtime, q, kv, runID := makeSecureBusRuntimeFixture(t, tool, securebus.DefaultPolicyConfig())
|
||||||
|
|
||||||
|
results, err := runtime.Execute(
|
||||||
|
fantasy.WithStepIndex(t.Context(), 0),
|
||||||
|
nil,
|
||||||
|
[]fantasy.ToolCallContent{{ToolCallID: "call-1", ToolName: "echo", Input: `{"text":`}},
|
||||||
|
nil,
|
||||||
|
)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, results, 1)
|
||||||
|
|
||||||
|
errorResult, ok := results[0].Result.(fantasy.ToolResultOutputContentError)
|
||||||
|
require.True(t, ok)
|
||||||
|
assert.Contains(t, errorResult.Error.Error(), "policy violation")
|
||||||
|
assert.Equal(t, int32(0), tool.calls.Load())
|
||||||
|
|
||||||
|
rows, err := q.ListAgentToolResultsByRunID(t.Context(), sqlc.ListAgentToolResultsByRunIDParams{RunID: runID, Lim: 10})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, rows, 1)
|
||||||
|
require.NotNil(t, rows[0].Preview)
|
||||||
|
assert.Contains(t, *rows[0].Preview, "policy violation")
|
||||||
|
|
||||||
|
persisted, err := kv.Get(t.Context(), rows[0].FullKey)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Contains(t, string(persisted), `"type":"error"`)
|
||||||
|
assert.Contains(t, string(persisted), "policy violation")
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,7 @@
|
||||||
package agent
|
package agent
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
@ -11,6 +12,7 @@ import (
|
||||||
"github.com/ZanzyTHEbar/dragonscale/pkg/bus"
|
"github.com/ZanzyTHEbar/dragonscale/pkg/bus"
|
||||||
"github.com/ZanzyTHEbar/dragonscale/pkg/config"
|
"github.com/ZanzyTHEbar/dragonscale/pkg/config"
|
||||||
memstore "github.com/ZanzyTHEbar/dragonscale/pkg/memory/store"
|
memstore "github.com/ZanzyTHEbar/dragonscale/pkg/memory/store"
|
||||||
|
"github.com/ZanzyTHEbar/dragonscale/pkg/tools"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
)
|
)
|
||||||
|
|
@ -109,3 +111,31 @@ func invokeMemoryAction(t *testing.T, tool *MemGPTTool, args map[string]interfac
|
||||||
require.True(t, response.Success)
|
require.True(t, response.Success)
|
||||||
return response
|
return response
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestMemGPTTool_PrefersContextSessionOverActiveSession(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
al := newSessionBoundAgentLoop(t)
|
||||||
|
toolAny, ok := al.tools.Get("memory")
|
||||||
|
require.True(t, ok)
|
||||||
|
memTool, ok := toolAny.(*MemGPTTool)
|
||||||
|
require.True(t, ok)
|
||||||
|
|
||||||
|
al.activeSessionKey.Store("session-b")
|
||||||
|
ctx := tools.WithSessionKey(context.Background(), "session-a")
|
||||||
|
writeA := memTool.Execute(ctx, map[string]interface{}{
|
||||||
|
"action": "write",
|
||||||
|
"content": "context-bound memory",
|
||||||
|
"tier": "recall",
|
||||||
|
"sector": "semantic",
|
||||||
|
})
|
||||||
|
require.False(t, writeA.IsError)
|
||||||
|
|
||||||
|
sessionAItems, err := al.memDelegate.ListRecallItems(t.Context(), pkgroot.NAME, "session-a", 10, 0)
|
||||||
|
require.NoError(t, err)
|
||||||
|
sessionBItems, err := al.memDelegate.ListRecallItems(t.Context(), pkgroot.NAME, "session-b", 10, 0)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotEmpty(t, sessionAItems)
|
||||||
|
assert.Equal(t, "context-bound memory", sessionAItems[0].Content)
|
||||||
|
assert.Empty(t, sessionBItems)
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -24,6 +24,7 @@ func (al *AgentLoop) endTask(ctx context.Context, conversationID, runID ids.UUID
|
||||||
if al.memDelegate == nil {
|
if al.memDelegate == nil {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
ctx = context.WithoutCancel(ctx)
|
||||||
|
|
||||||
// Store self-report scores if the delegate implements RLStore
|
// Store self-report scores if the delegate implements RLStore
|
||||||
if rlStore, ok := al.memDelegate.(interface {
|
if rlStore, ok := al.memDelegate.(interface {
|
||||||
|
|
|
||||||
|
|
@ -30,21 +30,47 @@ func MakeUnifiedRunLoopFunc(al *AgentLoop) tools.RunLoopFunc {
|
||||||
return nil, fmt.Errorf("unified runtime dependencies are not initialized")
|
return nil, fmt.Errorf("unified runtime dependencies are not initialized")
|
||||||
}
|
}
|
||||||
|
|
||||||
baseSession := fmt.Sprintf("%s:%s", channel, chatID)
|
baseSession := strings.TrimSpace(tools.SessionKeyFromContext(ctx))
|
||||||
if v := al.activeSessionKey.Load(); v != nil {
|
if baseSession == "" {
|
||||||
if active, ok := v.(string); ok && strings.TrimSpace(active) != "" {
|
baseSession = strings.TrimSpace(tools.DelegationSessionKeyFromContext(ctx))
|
||||||
baseSession = active
|
}
|
||||||
}
|
if baseSession == "" {
|
||||||
|
baseSession = fmt.Sprintf("%s:%s", channel, chatID)
|
||||||
}
|
}
|
||||||
if strings.TrimSpace(baseSession) == "" {
|
if strings.TrimSpace(baseSession) == "" {
|
||||||
baseSession = "subagent:default"
|
baseSession = "subagent:default"
|
||||||
}
|
}
|
||||||
sessionKey := fmt.Sprintf("%s::subagent::%s", baseSession, ids.New().String()[:8])
|
sessionKey := fmt.Sprintf("%s::subagent::%s", baseSession, ids.New().String()[:8])
|
||||||
|
ctx = tools.WithSessionKey(withToolSessionKey(ctx, sessionKey), sessionKey)
|
||||||
|
|
||||||
conversationID, runID, err := al.prepareRuntimeState(ctx, sessionKey)
|
conversationID, runID, err := al.prepareRuntimeState(ctx, sessionKey)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
opts := processOptions{
|
||||||
|
SessionKey: sessionKey,
|
||||||
|
Channel: channel,
|
||||||
|
ChatID: chatID,
|
||||||
|
UserMessage: userPrompt,
|
||||||
|
ConversationID: conversationID,
|
||||||
|
RunID: runID,
|
||||||
|
}
|
||||||
|
metrics := agentRunMetrics{}
|
||||||
|
defer func() {
|
||||||
|
if err != nil {
|
||||||
|
al.persistFailedRun(ctx, opts, err)
|
||||||
|
if endErr := al.endTask(ctx, conversationID, runID, TaskCompletion{
|
||||||
|
TaskID: sessionKey,
|
||||||
|
Description: userPrompt,
|
||||||
|
Completed: false,
|
||||||
|
}); endErr != nil {
|
||||||
|
logger.WarnCF("toolloop", "Failed to record failed subagent task completion", map[string]any{"error": endErr.Error(), "session": sessionKey})
|
||||||
|
}
|
||||||
|
_ = al.sessions.Save(sessionKey)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
al.persistRunCheckpoint(ctx, opts, metrics)
|
||||||
|
}()
|
||||||
|
|
||||||
baseRuntime := OffloadingToolRuntime{
|
baseRuntime := OffloadingToolRuntime{
|
||||||
Base: fantasy.DAGToolRuntime{MaxConcurrency: defaultToolMaxConcurrency},
|
Base: fantasy.DAGToolRuntime{MaxConcurrency: defaultToolMaxConcurrency},
|
||||||
|
|
@ -54,27 +80,39 @@ func MakeUnifiedRunLoopFunc(al *AgentLoop) tools.RunLoopFunc {
|
||||||
RunID: runID,
|
RunID: runID,
|
||||||
ThresholdChars: al.offloadThresholdChars,
|
ThresholdChars: al.offloadThresholdChars,
|
||||||
}
|
}
|
||||||
toolRuntime := SecureBusToolRuntime{
|
|
||||||
Base: baseRuntime,
|
|
||||||
Bus: al.secureBus,
|
|
||||||
SessionKey: sessionKey,
|
|
||||||
UserPrompt: userPrompt,
|
|
||||||
StateStore: al.stateStore,
|
|
||||||
RunID: runID,
|
|
||||||
}
|
|
||||||
|
|
||||||
extraTools := make([]fantasy.AgentTool, 0, 1)
|
extraTools := make([]fantasy.AgentTool, 0, 1)
|
||||||
if al.toolResultSearch != nil {
|
if al.toolResultSearch != nil {
|
||||||
extraTools = append(extraTools, al.toolResultSearch)
|
extraTools = append(extraTools, al.toolResultSearch)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
toolRuntime := SecureBusToolRuntime{
|
||||||
|
Offloader: baseRuntime,
|
||||||
|
FantasyTools: fantasyToolMap(extraTools, al.tools),
|
||||||
|
Bus: al.secureBus,
|
||||||
|
SessionKey: sessionKey,
|
||||||
|
Channel: channel,
|
||||||
|
ChatID: chatID,
|
||||||
|
UserPrompt: userPrompt,
|
||||||
|
StateStore: al.stateStore,
|
||||||
|
RunID: runID,
|
||||||
|
}
|
||||||
|
|
||||||
al.sessions.AddMessage(sessionKey, "user", userPrompt)
|
al.sessions.AddMessage(sessionKey, "user", userPrompt)
|
||||||
result, err := runToolLoopWithRuntime(ctx, config, systemPrompt, userPrompt, channel, chatID, al.memoryStore, sessionKey, toolRuntime, extraTools)
|
result, err := runToolLoopWithRuntime(ctx, config, systemPrompt, userPrompt, channel, chatID, al.memoryStore, sessionKey, toolRuntime, extraTools)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
metrics.StepCount = result.Iterations
|
||||||
al.sessions.AddMessage(sessionKey, "assistant", result.Content)
|
al.sessions.AddMessage(sessionKey, "assistant", result.Content)
|
||||||
al.sessions.Save(sessionKey)
|
al.sessions.Save(sessionKey)
|
||||||
|
if endErr := al.endTask(ctx, conversationID, runID, TaskCompletion{
|
||||||
|
TaskID: sessionKey,
|
||||||
|
Description: userPrompt,
|
||||||
|
Completed: true,
|
||||||
|
}); endErr != nil {
|
||||||
|
logger.WarnCF("toolloop", "Failed to record subagent task completion", map[string]any{"error": endErr.Error(), "session": sessionKey})
|
||||||
|
}
|
||||||
al.maybeSummarize(ctx, sessionKey, channel, chatID)
|
al.maybeSummarize(ctx, sessionKey, channel, chatID)
|
||||||
return result, nil
|
return result, nil
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -90,26 +90,19 @@ func (a *DragonToolAdapter) Run(ctx context.Context, call fantasy.ToolCall) (fan
|
||||||
return fantasy.NewTextErrorResponse(fmt.Sprintf("invalid arguments: %v", err)), nil
|
return fantasy.NewTextErrorResponse(fmt.Sprintf("invalid arguments: %v", err)), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// 2. Set context for ContextualTool implementations.
|
execCtx := tools.WithExecutionTarget(ctx, a.channel, a.chatID)
|
||||||
if ct, ok := a.inner.(tools.ContextualTool); ok {
|
execCtx = tools.WithAsyncCallback(execCtx, func(_ context.Context, result *tools.ToolResult) {
|
||||||
ct.SetContext(a.channel, a.chatID)
|
if result != nil && result.ForUser != "" && !result.Silent && a.bus != nil {
|
||||||
}
|
a.bus.PublishOutbound(bus.OutboundMessage{
|
||||||
|
Channel: a.channel,
|
||||||
// 3. Wire async callback for AsyncTool implementations.
|
ChatID: a.chatID,
|
||||||
if at, ok := a.inner.(tools.AsyncTool); ok {
|
Content: result.ForUser,
|
||||||
at.SetCallback(func(_ context.Context, result *tools.ToolResult) {
|
})
|
||||||
if result != nil && result.ForUser != "" && !result.Silent && a.bus != nil {
|
}
|
||||||
a.bus.PublishOutbound(bus.OutboundMessage{
|
})
|
||||||
Channel: a.channel,
|
|
||||||
ChatID: a.chatID,
|
|
||||||
Content: result.ForUser,
|
|
||||||
})
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// 4. Execute the DragonScale tool.
|
// 4. Execute the DragonScale tool.
|
||||||
result := a.inner.Execute(ctx, args)
|
result := a.inner.Execute(execCtx, args)
|
||||||
if result == nil {
|
if result == nil {
|
||||||
return fantasy.NewTextErrorResponse("tool returned nil result"), nil
|
return fantasy.NewTextErrorResponse("tool returned nil result"), nil
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -78,7 +78,11 @@ func (t *mockContextualTool) Description() string { return "Contextual tool" }
|
||||||
func (t *mockContextualTool) Parameters() map[string]interface{} {
|
func (t *mockContextualTool) Parameters() map[string]interface{} {
|
||||||
return map[string]interface{}{"type": "object", "properties": map[string]interface{}{}}
|
return map[string]interface{}{"type": "object", "properties": map[string]interface{}{}}
|
||||||
}
|
}
|
||||||
func (t *mockContextualTool) Execute(_ context.Context, _ map[string]interface{}) *tools.ToolResult {
|
func (t *mockContextualTool) Execute(ctx context.Context, _ map[string]interface{}) *tools.ToolResult {
|
||||||
|
if channel, chatID := tools.ExecutionTargetFromContext(ctx); channel != "" || chatID != "" {
|
||||||
|
t.channel = channel
|
||||||
|
t.chatID = chatID
|
||||||
|
}
|
||||||
return &tools.ToolResult{
|
return &tools.ToolResult{
|
||||||
ForLLM: "channel=" + t.channel + " chat=" + t.chatID,
|
ForLLM: "channel=" + t.channel + " chat=" + t.chatID,
|
||||||
Silent: true,
|
Silent: true,
|
||||||
|
|
|
||||||
|
|
@ -448,8 +448,8 @@ type Querier interface {
|
||||||
// created_at
|
// created_at
|
||||||
// FROM task_completions
|
// FROM task_completions
|
||||||
// WHERE agent_id = ?1
|
// WHERE agent_id = ?1
|
||||||
// AND created_at > ?2
|
// AND unixepoch(created_at) > unixepoch(?2)
|
||||||
// AND completed = 1
|
// AND completed = 1
|
||||||
// ORDER BY created_at ASC
|
// ORDER BY created_at ASC
|
||||||
GetCompletedTasks(ctx context.Context, arg GetCompletedTasksParams) ([]TaskCompletion, error)
|
GetCompletedTasks(ctx context.Context, arg GetCompletedTasksParams) ([]TaskCompletion, error)
|
||||||
//GetDAGNodeBySnapshotAndNodeID
|
//GetDAGNodeBySnapshotAndNodeID
|
||||||
|
|
@ -1154,7 +1154,7 @@ type Querier interface {
|
||||||
//
|
//
|
||||||
// SELECT DISTINCT agent_id
|
// SELECT DISTINCT agent_id
|
||||||
// FROM task_completions
|
// FROM task_completions
|
||||||
// WHERE created_at > ?1
|
// WHERE unixepoch(created_at) > unixepoch(?1)
|
||||||
// ORDER BY agent_id
|
// ORDER BY agent_id
|
||||||
ListActiveAgents(ctx context.Context, arg ListActiveAgentsParams) ([]string, error)
|
ListActiveAgents(ctx context.Context, arg ListActiveAgentsParams) ([]string, error)
|
||||||
//ListAgentCheckpointsByConversationID
|
//ListAgentCheckpointsByConversationID
|
||||||
|
|
|
||||||
|
|
@ -171,8 +171,8 @@ SELECT id,
|
||||||
created_at
|
created_at
|
||||||
FROM task_completions
|
FROM task_completions
|
||||||
WHERE agent_id = sqlc.arg(agent_id)
|
WHERE agent_id = sqlc.arg(agent_id)
|
||||||
AND created_at > sqlc.arg(since)
|
AND unixepoch(created_at) > unixepoch(sqlc.arg(since))
|
||||||
AND completed = 1
|
AND completed = 1
|
||||||
ORDER BY created_at ASC;
|
ORDER BY created_at ASC;
|
||||||
-- name: StoreTaskRetrieval :exec
|
-- name: StoreTaskRetrieval :exec
|
||||||
-- Store a memory retrieval record for a task
|
-- Store a memory retrieval record for a task
|
||||||
|
|
@ -197,7 +197,7 @@ WHERE tr.task_id = sqlc.arg(task_id);
|
||||||
-- Get all unique agent IDs that have completed tasks (for multi-agent processing)
|
-- Get all unique agent IDs that have completed tasks (for multi-agent processing)
|
||||||
SELECT DISTINCT agent_id
|
SELECT DISTINCT agent_id
|
||||||
FROM task_completions
|
FROM task_completions
|
||||||
WHERE created_at > sqlc.arg(since)
|
WHERE unixepoch(created_at) > unixepoch(sqlc.arg(since))
|
||||||
ORDER BY agent_id;
|
ORDER BY agent_id;
|
||||||
-- name: GetHighTokenSessions :many
|
-- name: GetHighTokenSessions :many
|
||||||
-- Get sessions with high token usage grouped by conversation/agent
|
-- Get sessions with high token usage grouped by conversation/agent
|
||||||
|
|
|
||||||
|
|
@ -27,14 +27,14 @@ SELECT id,
|
||||||
created_at
|
created_at
|
||||||
FROM task_completions
|
FROM task_completions
|
||||||
WHERE agent_id = ?1
|
WHERE agent_id = ?1
|
||||||
AND created_at > ?2
|
AND unixepoch(created_at) > unixepoch(?2)
|
||||||
AND completed = 1
|
AND completed = 1
|
||||||
ORDER BY created_at ASC
|
ORDER BY created_at ASC
|
||||||
`
|
`
|
||||||
|
|
||||||
type GetCompletedTasksParams struct {
|
type GetCompletedTasksParams struct {
|
||||||
AgentID string `db:"agent_id" json:"agent_id"`
|
AgentID string `db:"agent_id" json:"agent_id"`
|
||||||
Since time.Time `db:"since" json:"since"`
|
Since interface{} `db:"since" json:"since"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// Get tasks completed since the given time for RL processing
|
// Get tasks completed since the given time for RL processing
|
||||||
|
|
@ -52,8 +52,8 @@ type GetCompletedTasksParams struct {
|
||||||
// created_at
|
// created_at
|
||||||
// FROM task_completions
|
// FROM task_completions
|
||||||
// WHERE agent_id = ?1
|
// WHERE agent_id = ?1
|
||||||
// AND created_at > ?2
|
// AND unixepoch(created_at) > unixepoch(?2)
|
||||||
// AND completed = 1
|
// AND completed = 1
|
||||||
// ORDER BY created_at ASC
|
// ORDER BY created_at ASC
|
||||||
func (q *Queries) GetCompletedTasks(ctx context.Context, arg GetCompletedTasksParams) ([]TaskCompletion, error) {
|
func (q *Queries) GetCompletedTasks(ctx context.Context, arg GetCompletedTasksParams) ([]TaskCompletion, error) {
|
||||||
rows, err := q.db.QueryContext(ctx, GetCompletedTasks, arg.AgentID, arg.Since)
|
rows, err := q.db.QueryContext(ctx, GetCompletedTasks, arg.AgentID, arg.Since)
|
||||||
|
|
@ -398,19 +398,19 @@ func (q *Queries) IncrementTaskRetrievalCount(ctx context.Context, arg Increment
|
||||||
const ListActiveAgents = `-- name: ListActiveAgents :many
|
const ListActiveAgents = `-- name: ListActiveAgents :many
|
||||||
SELECT DISTINCT agent_id
|
SELECT DISTINCT agent_id
|
||||||
FROM task_completions
|
FROM task_completions
|
||||||
WHERE created_at > ?1
|
WHERE unixepoch(created_at) > unixepoch(?1)
|
||||||
ORDER BY agent_id
|
ORDER BY agent_id
|
||||||
`
|
`
|
||||||
|
|
||||||
type ListActiveAgentsParams struct {
|
type ListActiveAgentsParams struct {
|
||||||
Since time.Time `db:"since" json:"since"`
|
Since interface{} `db:"since" json:"since"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// Get all unique agent IDs that have completed tasks (for multi-agent processing)
|
// Get all unique agent IDs that have completed tasks (for multi-agent processing)
|
||||||
//
|
//
|
||||||
// SELECT DISTINCT agent_id
|
// SELECT DISTINCT agent_id
|
||||||
// FROM task_completions
|
// FROM task_completions
|
||||||
// WHERE created_at > ?1
|
// WHERE unixepoch(created_at) > unixepoch(?1)
|
||||||
// ORDER BY agent_id
|
// ORDER BY agent_id
|
||||||
func (q *Queries) ListActiveAgents(ctx context.Context, arg ListActiveAgentsParams) ([]string, error) {
|
func (q *Queries) ListActiveAgents(ctx context.Context, arg ListActiveAgentsParams) ([]string, error) {
|
||||||
rows, err := q.db.QueryContext(ctx, ListActiveAgents, arg.Since)
|
rows, err := q.db.QueryContext(ctx, ListActiveAgents, arg.Since)
|
||||||
|
|
|
||||||
|
|
@ -24,6 +24,12 @@ type State struct {
|
||||||
// LastChatID is the last chat ID used for communication
|
// LastChatID is the last chat ID used for communication
|
||||||
LastChatID string `json:"last_chat_id,omitzero"`
|
LastChatID string `json:"last_chat_id,omitzero"`
|
||||||
|
|
||||||
|
// LastSessionKey is the last non-ephemeral user session processed.
|
||||||
|
LastSessionKey string `json:"last_session_key,omitzero"`
|
||||||
|
|
||||||
|
// LastSessionKeysByTarget tracks the last durable session key per channel/chat target.
|
||||||
|
LastSessionKeysByTarget map[string]string `json:"last_session_keys_by_target,omitzero"`
|
||||||
|
|
||||||
// Timestamp is the last time this state was updated
|
// Timestamp is the last time this state was updated
|
||||||
Timestamp time.Time `json:"timestamp"`
|
Timestamp time.Time `json:"timestamp"`
|
||||||
}
|
}
|
||||||
|
|
@ -107,6 +113,37 @@ func (sm *Manager) SetChannelAndChatID(ctx context.Context, channel, chatID stri
|
||||||
return sm.persist(ctx)
|
return sm.persist(ctx)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SetLastSessionKey updates the last durable user session key.
|
||||||
|
func (sm *Manager) SetLastSessionKey(ctx context.Context, sessionKey string) error {
|
||||||
|
sm.mu.Lock()
|
||||||
|
defer sm.mu.Unlock()
|
||||||
|
|
||||||
|
sm.state.LastSessionKey = sessionKey
|
||||||
|
sm.state.Timestamp = time.Now()
|
||||||
|
|
||||||
|
return sm.persist(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetLastSessionKeyForTarget updates the last durable session key for a specific channel/chat target.
|
||||||
|
func (sm *Manager) SetLastSessionKeyForTarget(ctx context.Context, channel, chatID, sessionKey string) error {
|
||||||
|
sm.mu.Lock()
|
||||||
|
defer sm.mu.Unlock()
|
||||||
|
|
||||||
|
if sm.state.LastSessionKeysByTarget == nil {
|
||||||
|
sm.state.LastSessionKeysByTarget = make(map[string]string)
|
||||||
|
}
|
||||||
|
key := channel + ":" + chatID
|
||||||
|
if sessionKey == "" {
|
||||||
|
delete(sm.state.LastSessionKeysByTarget, key)
|
||||||
|
} else {
|
||||||
|
sm.state.LastSessionKeysByTarget[key] = sessionKey
|
||||||
|
}
|
||||||
|
sm.state.LastSessionKey = sessionKey
|
||||||
|
sm.state.Timestamp = time.Now()
|
||||||
|
|
||||||
|
return sm.persist(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
// persist writes the current state to the delegate (KV) if available.
|
// persist writes the current state to the delegate (KV) if available.
|
||||||
// Must be called with the lock held.
|
// Must be called with the lock held.
|
||||||
func (sm *Manager) persist(ctx context.Context) error {
|
func (sm *Manager) persist(ctx context.Context) error {
|
||||||
|
|
@ -141,6 +178,23 @@ func (sm *Manager) GetLastChatID() string {
|
||||||
return sm.state.LastChatID
|
return sm.state.LastChatID
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// GetLastSessionKey returns the last durable session key from the state.
|
||||||
|
func (sm *Manager) GetLastSessionKey() string {
|
||||||
|
sm.mu.RLock()
|
||||||
|
defer sm.mu.RUnlock()
|
||||||
|
return sm.state.LastSessionKey
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetLastSessionKeyForTarget returns the last durable session key for a specific channel/chat target.
|
||||||
|
func (sm *Manager) GetLastSessionKeyForTarget(channel, chatID string) string {
|
||||||
|
sm.mu.RLock()
|
||||||
|
defer sm.mu.RUnlock()
|
||||||
|
if sm.state.LastSessionKeysByTarget == nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return sm.state.LastSessionKeysByTarget[channel+":"+chatID]
|
||||||
|
}
|
||||||
|
|
||||||
// GetTimestamp returns the timestamp of the last state update.
|
// GetTimestamp returns the timestamp of the last state update.
|
||||||
func (sm *Manager) GetTimestamp() time.Time {
|
func (sm *Manager) GetTimestamp() time.Time {
|
||||||
sm.mu.RLock()
|
sm.mu.RLock()
|
||||||
|
|
|
||||||
|
|
@ -112,3 +112,31 @@ func TestStateStruct(t *testing.T) {
|
||||||
t.Errorf("Expected LastChatID 'test-chat-id', got '%s'", state.LastChatID)
|
t.Errorf("Expected LastChatID 'test-chat-id', got '%s'", state.LastChatID)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestSetLastSessionKeyForTarget(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
sm := NewManager("/tmp/test")
|
||||||
|
err := sm.SetLastSessionKeyForTarget(context.Background(), "telegram", "chat-1", "session-a")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("SetLastSessionKeyForTarget failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if got := sm.GetLastSessionKeyForTarget("telegram", "chat-1"); got != "session-a" {
|
||||||
|
t.Fatalf("expected session-a for target, got %q", got)
|
||||||
|
}
|
||||||
|
if got := sm.GetLastSessionKey(); got != "session-a" {
|
||||||
|
t.Fatalf("expected global last session to mirror latest update, got %q", got)
|
||||||
|
}
|
||||||
|
if got := sm.GetLastSessionKeyForTarget("telegram", "chat-2"); got != "" {
|
||||||
|
t.Fatalf("expected empty session for other target, got %q", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
err = sm.SetLastSessionKeyForTarget(context.Background(), "telegram", "chat-1", "")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("clearing SetLastSessionKeyForTarget failed: %v", err)
|
||||||
|
}
|
||||||
|
if got := sm.GetLastSessionKeyForTarget("telegram", "chat-1"); got != "" {
|
||||||
|
t.Fatalf("expected cleared target session, got %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -92,6 +92,7 @@ func (t *AgenticMapTool) Execute(ctx context.Context, args map[string]interface{
|
||||||
if t.manager == nil && t.runtime == nil {
|
if t.manager == nil && t.runtime == nil {
|
||||||
return ErrorResult("agentic_map manager is not configured").WithError(fmt.Errorf("agentic_map manager is nil"))
|
return ErrorResult("agentic_map manager is not configured").WithError(fmt.Errorf("agentic_map manager is nil"))
|
||||||
}
|
}
|
||||||
|
originChannel, originChatID := ResolveExecutionTarget(ctx, t.originChannel, t.originChatID)
|
||||||
|
|
||||||
items, usedJSONL, err := parseMapBoundaryItems(args)
|
items, usedJSONL, err := parseMapBoundaryItems(args)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|
@ -142,8 +143,8 @@ func (t *AgenticMapTool) Execute(ctx context.Context, args map[string]interface{
|
||||||
MaxRetries: uint16(maxRetries),
|
MaxRetries: uint16(maxRetries),
|
||||||
DelegatedScope: delegatedScope,
|
DelegatedScope: delegatedScope,
|
||||||
KeptWork: keptWork,
|
KeptWork: keptWork,
|
||||||
OriginChannel: t.originChannel,
|
OriginChannel: originChannel,
|
||||||
OriginChatID: t.originChatID,
|
OriginChatID: originChatID,
|
||||||
}, items, idempotencyKey)
|
}, items, idempotencyKey)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return ErrorResult(fmt.Sprintf("failed to enqueue agentic_map run: %v", err)).WithError(err)
|
return ErrorResult(fmt.Sprintf("failed to enqueue agentic_map run: %v", err)).WithError(err)
|
||||||
|
|
@ -189,7 +190,7 @@ func (t *AgenticMapTool) Execute(ctx context.Context, args map[string]interface{
|
||||||
}
|
}
|
||||||
|
|
||||||
subTool := NewSubagentTool(t.manager)
|
subTool := NewSubagentTool(t.manager)
|
||||||
subTool.SetContext(t.originChannel, t.originChatID)
|
subTool.SetContext(originChannel, originChatID)
|
||||||
|
|
||||||
type itemResult struct {
|
type itemResult struct {
|
||||||
Index int `json:"index"`
|
Index int `json:"index"`
|
||||||
|
|
|
||||||
|
|
@ -124,7 +124,8 @@ func (t *ToolCallTool) Execute(ctx context.Context, args map[string]interface{})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return t.registry.ExecuteWithContext(ctx, toolName, toolArgs, t.channel, t.chatID, nil)
|
channel, chatID := ResolveExecutionTarget(ctx, t.channel, t.chatID)
|
||||||
|
return t.registry.ExecuteWithContext(ctx, toolName, toolArgs, channel, chatID, AsyncCallbackFromContext(ctx))
|
||||||
}
|
}
|
||||||
|
|
||||||
// schemaHintError returns an error result that includes the tool's expected
|
// schemaHintError returns an error result that includes the tool's expected
|
||||||
|
|
|
||||||
|
|
@ -2,8 +2,10 @@ package tools
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"fmt"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
"time"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestToolCallTool_Name(t *testing.T) {
|
func TestToolCallTool_Name(t *testing.T) {
|
||||||
|
|
@ -202,9 +204,9 @@ func TestToolCallTool_ContextPropagation(t *testing.T) {
|
||||||
r.Register(ct)
|
r.Register(ct)
|
||||||
|
|
||||||
tc := NewToolCallTool(r)
|
tc := NewToolCallTool(r)
|
||||||
tc.SetContext("test-channel", "test-chat")
|
ctx := WithExecutionTarget(t.Context(), "test-channel", "test-chat")
|
||||||
|
|
||||||
tc.Execute(t.Context(), map[string]interface{}{
|
tc.Execute(ctx, map[string]interface{}{
|
||||||
"tool_name": "capture",
|
"tool_name": "capture",
|
||||||
"arguments": map[string]interface{}{},
|
"arguments": map[string]interface{}{},
|
||||||
})
|
})
|
||||||
|
|
@ -215,6 +217,82 @@ func TestToolCallTool_ContextPropagation(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestToolCallTool_ForwardsAsyncCallbackAndExecutionTarget(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
r := NewToolRegistry()
|
||||||
|
asyncTool := &callbackCaptureTool{}
|
||||||
|
r.Register(asyncTool)
|
||||||
|
tc := NewToolCallTool(r)
|
||||||
|
|
||||||
|
ctx := WithExecutionTarget(t.Context(), "telegram", "chat-77")
|
||||||
|
callbackDone := make(chan *ToolResult, 1)
|
||||||
|
ctx = WithAsyncCallback(ctx, func(_ context.Context, result *ToolResult) {
|
||||||
|
callbackDone <- result
|
||||||
|
})
|
||||||
|
|
||||||
|
result := tc.Execute(ctx, map[string]interface{}{
|
||||||
|
"tool_name": "spawn",
|
||||||
|
"arguments": map[string]interface{}{"task": "background"},
|
||||||
|
})
|
||||||
|
|
||||||
|
if result.IsError {
|
||||||
|
t.Fatalf("unexpected error: %s", result.ForLLM)
|
||||||
|
}
|
||||||
|
if !result.Async {
|
||||||
|
t.Fatal("expected async result from nested spawn tool")
|
||||||
|
}
|
||||||
|
|
||||||
|
select {
|
||||||
|
case callbackResult := <-callbackDone:
|
||||||
|
if callbackResult == nil {
|
||||||
|
t.Fatal("expected callback result")
|
||||||
|
}
|
||||||
|
if callbackResult.ForUser != "async completion on telegram:chat-77" {
|
||||||
|
t.Fatalf("unexpected callback result: %s", callbackResult.ForUser)
|
||||||
|
}
|
||||||
|
case <-time.After(2 * time.Second):
|
||||||
|
t.Fatal("timed out waiting for async callback")
|
||||||
|
}
|
||||||
|
|
||||||
|
if asyncTool.lastChannel != "telegram" || asyncTool.lastChatID != "chat-77" {
|
||||||
|
t.Fatalf("expected execution target propagation, got %s:%s", asyncTool.lastChannel, asyncTool.lastChatID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestToolCallTool_LegacyAsyncToolStillReceivesCallback(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
r := NewToolRegistry()
|
||||||
|
legacyTool := &legacyAsyncCallbackTool{}
|
||||||
|
r.Register(legacyTool)
|
||||||
|
tc := NewToolCallTool(r)
|
||||||
|
|
||||||
|
callbackDone := make(chan *ToolResult, 1)
|
||||||
|
ctx := WithAsyncCallback(t.Context(), func(_ context.Context, result *ToolResult) {
|
||||||
|
callbackDone <- result
|
||||||
|
})
|
||||||
|
|
||||||
|
result := tc.Execute(ctx, map[string]interface{}{
|
||||||
|
"tool_name": "legacy_async",
|
||||||
|
"arguments": map[string]interface{}{},
|
||||||
|
})
|
||||||
|
|
||||||
|
if result.IsError {
|
||||||
|
t.Fatalf("unexpected error: %s", result.ForLLM)
|
||||||
|
}
|
||||||
|
if !result.Async {
|
||||||
|
t.Fatal("expected async result from legacy async tool")
|
||||||
|
}
|
||||||
|
|
||||||
|
select {
|
||||||
|
case callbackResult := <-callbackDone:
|
||||||
|
if callbackResult == nil || callbackResult.ForUser != "legacy completion" {
|
||||||
|
t.Fatalf("unexpected legacy callback result: %#v", callbackResult)
|
||||||
|
}
|
||||||
|
case <-time.After(2 * time.Second):
|
||||||
|
t.Fatal("timed out waiting for legacy async callback")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestToolCallTool_ResourceProvider_LoadsResources(t *testing.T) {
|
func TestToolCallTool_ResourceProvider_LoadsResources(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
r := NewToolRegistry()
|
r := NewToolRegistry()
|
||||||
|
|
@ -361,6 +439,45 @@ func (c *contextCaptureTool) SetContext(channel, chatID string) {
|
||||||
c.lastChannel = channel
|
c.lastChannel = channel
|
||||||
c.lastChatID = chatID
|
c.lastChatID = chatID
|
||||||
}
|
}
|
||||||
func (c *contextCaptureTool) Execute(_ context.Context, _ map[string]interface{}) *ToolResult {
|
func (c *contextCaptureTool) Execute(ctx context.Context, _ map[string]interface{}) *ToolResult {
|
||||||
|
if channel, chatID := ExecutionTargetFromContext(ctx); channel != "" || chatID != "" {
|
||||||
|
c.lastChannel = channel
|
||||||
|
c.lastChatID = chatID
|
||||||
|
}
|
||||||
return &ToolResult{ForLLM: "captured"}
|
return &ToolResult{ForLLM: "captured"}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type callbackCaptureTool struct {
|
||||||
|
lastChannel string
|
||||||
|
lastChatID string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *callbackCaptureTool) Name() string { return "spawn" }
|
||||||
|
func (c *callbackCaptureTool) Description() string { return "captures async callback propagation" }
|
||||||
|
func (c *callbackCaptureTool) Parameters() map[string]interface{} {
|
||||||
|
return map[string]interface{}{"type": "object", "properties": map[string]interface{}{}}
|
||||||
|
}
|
||||||
|
func (c *callbackCaptureTool) Execute(ctx context.Context, _ map[string]interface{}) *ToolResult {
|
||||||
|
c.lastChannel, c.lastChatID = ExecutionTargetFromContext(ctx)
|
||||||
|
if callback := AsyncCallbackFromContext(ctx); callback != nil {
|
||||||
|
callback(ctx, &ToolResult{ForLLM: "done", ForUser: fmt.Sprintf("async completion on %s:%s", c.lastChannel, c.lastChatID)})
|
||||||
|
}
|
||||||
|
return AsyncResult("spawned")
|
||||||
|
}
|
||||||
|
|
||||||
|
type legacyAsyncCallbackTool struct {
|
||||||
|
callback AsyncCallback
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *legacyAsyncCallbackTool) Name() string { return "legacy_async" }
|
||||||
|
func (t *legacyAsyncCallbackTool) Description() string { return "legacy async callback tool" }
|
||||||
|
func (t *legacyAsyncCallbackTool) Parameters() map[string]interface{} {
|
||||||
|
return map[string]interface{}{"type": "object", "properties": map[string]interface{}{}}
|
||||||
|
}
|
||||||
|
func (t *legacyAsyncCallbackTool) SetCallback(cb AsyncCallback) { t.callback = cb }
|
||||||
|
func (t *legacyAsyncCallbackTool) Execute(ctx context.Context, _ map[string]interface{}) *ToolResult {
|
||||||
|
if t.callback != nil {
|
||||||
|
go t.callback(ctx, &ToolResult{ForLLM: "legacy completion", ForUser: "legacy completion"})
|
||||||
|
}
|
||||||
|
return AsyncResult("legacy started")
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -107,10 +107,11 @@ func (t *CronTool) Execute(ctx context.Context, args map[string]interface{}) *To
|
||||||
if !ok {
|
if !ok {
|
||||||
return ErrorResult("action is required")
|
return ErrorResult("action is required")
|
||||||
}
|
}
|
||||||
|
channel, chatID := ResolveExecutionTarget(ctx, t.channel, t.chatID)
|
||||||
|
|
||||||
switch action {
|
switch action {
|
||||||
case "add":
|
case "add":
|
||||||
return t.addJob(args)
|
return t.addJob(channel, chatID, args)
|
||||||
case "list":
|
case "list":
|
||||||
return t.listJobs()
|
return t.listJobs()
|
||||||
case "remove":
|
case "remove":
|
||||||
|
|
@ -124,12 +125,7 @@ func (t *CronTool) Execute(ctx context.Context, args map[string]interface{}) *To
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *CronTool) addJob(args map[string]interface{}) *ToolResult {
|
func (t *CronTool) addJob(channel, chatID string, args map[string]interface{}) *ToolResult {
|
||||||
t.mu.RLock()
|
|
||||||
channel := t.channel
|
|
||||||
chatID := t.chatID
|
|
||||||
t.mu.RUnlock()
|
|
||||||
|
|
||||||
if channel == "" || chatID == "" {
|
if channel == "" || chatID == "" {
|
||||||
return ErrorResult("no session context (channel/chat_id not set). Use this tool in an active conversation.")
|
return ErrorResult("no session context (channel/chat_id not set). Use this tool in an active conversation.")
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -130,7 +130,7 @@ func (t *DagExpandTool) Execute(ctx context.Context, args map[string]interface{}
|
||||||
|
|
||||||
sessionKey, _ := args["session_key"].(string)
|
sessionKey, _ := args["session_key"].(string)
|
||||||
if sessionKey == "" {
|
if sessionKey == "" {
|
||||||
sessionKey = t.deps.SessionFn()
|
sessionKey = ResolveSessionKey(ctx, t.deps.SessionFn)
|
||||||
}
|
}
|
||||||
if sessionKey == "" {
|
if sessionKey == "" {
|
||||||
sessionKey = "default"
|
sessionKey = "default"
|
||||||
|
|
@ -261,7 +261,7 @@ func (t *DagDescribeTool) Execute(ctx context.Context, args map[string]interface
|
||||||
|
|
||||||
sessionKey, _ := args["session_key"].(string)
|
sessionKey, _ := args["session_key"].(string)
|
||||||
if sessionKey == "" {
|
if sessionKey == "" {
|
||||||
sessionKey = t.deps.SessionFn()
|
sessionKey = ResolveSessionKey(ctx, t.deps.SessionFn)
|
||||||
}
|
}
|
||||||
if sessionKey == "" {
|
if sessionKey == "" {
|
||||||
sessionKey = "default"
|
sessionKey = "default"
|
||||||
|
|
@ -391,7 +391,7 @@ func (t *DagGrepTool) Execute(ctx context.Context, args map[string]interface{})
|
||||||
|
|
||||||
sessionKey, _ := args["session_key"].(string)
|
sessionKey, _ := args["session_key"].(string)
|
||||||
if sessionKey == "" {
|
if sessionKey == "" {
|
||||||
sessionKey = t.deps.SessionFn()
|
sessionKey = ResolveSessionKey(ctx, t.deps.SessionFn)
|
||||||
}
|
}
|
||||||
if sessionKey == "" {
|
if sessionKey == "" {
|
||||||
sessionKey = "default"
|
sessionKey = "default"
|
||||||
|
|
|
||||||
|
|
@ -140,7 +140,7 @@ type StartFocusTool struct {
|
||||||
delegate KVStore
|
delegate KVStore
|
||||||
sessions *session.SessionManager
|
sessions *session.SessionManager
|
||||||
sessionKey func() string
|
sessionKey func() string
|
||||||
OnChange func() // called after focus state changes; used for cache invalidation
|
OnChange func(sessionKey string) // called after focus state changes; used for cache invalidation
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewStartFocusTool(delegate KVStore, sessions *session.SessionManager, sessionKeyFn func() string) *StartFocusTool {
|
func NewStartFocusTool(delegate KVStore, sessions *session.SessionManager, sessionKeyFn func() string) *StartFocusTool {
|
||||||
|
|
@ -203,7 +203,7 @@ func (t *StartFocusTool) Execute(ctx context.Context, args map[string]interface{
|
||||||
return ErrorResult(fmt.Sprintf("deadline %v", err))
|
return ErrorResult(fmt.Sprintf("deadline %v", err))
|
||||||
}
|
}
|
||||||
|
|
||||||
sk := t.sessionKey()
|
sk := ResolveSessionKey(ctx, t.sessionKey)
|
||||||
if sk == "" {
|
if sk == "" {
|
||||||
return ErrorResult("no active session")
|
return ErrorResult("no active session")
|
||||||
}
|
}
|
||||||
|
|
@ -244,7 +244,7 @@ func (t *StartFocusTool) Execute(ctx context.Context, args map[string]interface{
|
||||||
}
|
}
|
||||||
|
|
||||||
if t.OnChange != nil {
|
if t.OnChange != nil {
|
||||||
t.OnChange()
|
t.OnChange(sk)
|
||||||
}
|
}
|
||||||
|
|
||||||
return SilentResult(fmt.Sprintf("Focus started on: %s\n%sCheckpoint at message %d. Explore freely, then call complete_focus when done.", topic, goalLine, state.CheckpointIndex))
|
return SilentResult(fmt.Sprintf("Focus started on: %s\n%sCheckpoint at message %d. Explore freely, then call complete_focus when done.", topic, goalLine, state.CheckpointIndex))
|
||||||
|
|
@ -255,7 +255,7 @@ type CompleteFocusTool struct {
|
||||||
delegate KVStore
|
delegate KVStore
|
||||||
sessions *session.SessionManager
|
sessions *session.SessionManager
|
||||||
sessionKey func() string
|
sessionKey func() string
|
||||||
OnChange func() // called after focus state + knowledge changes; used for cache invalidation
|
OnChange func(sessionKey string) // called after focus state + knowledge changes; used for cache invalidation
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewCompleteFocusTool(delegate KVStore, sessions *session.SessionManager, sessionKeyFn func() string) *CompleteFocusTool {
|
func NewCompleteFocusTool(delegate KVStore, sessions *session.SessionManager, sessionKeyFn func() string) *CompleteFocusTool {
|
||||||
|
|
@ -306,7 +306,7 @@ func (t *CompleteFocusTool) Execute(ctx context.Context, args map[string]interfa
|
||||||
return ErrorResult("summary is required")
|
return ErrorResult("summary is required")
|
||||||
}
|
}
|
||||||
|
|
||||||
sk := t.sessionKey()
|
sk := ResolveSessionKey(ctx, t.sessionKey)
|
||||||
if sk == "" {
|
if sk == "" {
|
||||||
return ErrorResult("no active session")
|
return ErrorResult("no active session")
|
||||||
}
|
}
|
||||||
|
|
@ -374,7 +374,7 @@ func (t *CompleteFocusTool) Execute(ctx context.Context, args map[string]interfa
|
||||||
})
|
})
|
||||||
|
|
||||||
if t.OnChange != nil {
|
if t.OnChange != nil {
|
||||||
t.OnChange()
|
t.OnChange(sk)
|
||||||
}
|
}
|
||||||
|
|
||||||
return SilentResult(fmt.Sprintf(
|
return SilentResult(fmt.Sprintf(
|
||||||
|
|
|
||||||
|
|
@ -50,7 +50,7 @@ func (t *FocusHistoryTool) Parameters() map[string]interface{} {
|
||||||
func (t *FocusHistoryTool) Execute(ctx context.Context, args map[string]interface{}) *ToolResult {
|
func (t *FocusHistoryTool) Execute(ctx context.Context, args map[string]interface{}) *ToolResult {
|
||||||
sessionKey := ""
|
sessionKey := ""
|
||||||
if t.sessionKey != nil {
|
if t.sessionKey != nil {
|
||||||
sessionKey = t.sessionKey()
|
sessionKey = ResolveSessionKey(ctx, t.sessionKey)
|
||||||
}
|
}
|
||||||
if strings.TrimSpace(sessionKey) == "" {
|
if strings.TrimSpace(sessionKey) == "" {
|
||||||
return ErrorResult("no active session")
|
return ErrorResult("no active session")
|
||||||
|
|
|
||||||
|
|
@ -11,7 +11,6 @@ type MessageTool struct {
|
||||||
sendCallback SendCallback
|
sendCallback SendCallback
|
||||||
defaultChannel string
|
defaultChannel string
|
||||||
defaultChatID string
|
defaultChatID string
|
||||||
sentInRound bool // Tracks whether a message was sent in the current processing round
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewMessageTool() *MessageTool {
|
func NewMessageTool() *MessageTool {
|
||||||
|
|
@ -50,12 +49,6 @@ func (t *MessageTool) Parameters() map[string]interface{} {
|
||||||
func (t *MessageTool) SetContext(channel, chatID string) {
|
func (t *MessageTool) SetContext(channel, chatID string) {
|
||||||
t.defaultChannel = channel
|
t.defaultChannel = channel
|
||||||
t.defaultChatID = chatID
|
t.defaultChatID = chatID
|
||||||
t.sentInRound = false // Reset send tracking for new processing round
|
|
||||||
}
|
|
||||||
|
|
||||||
// HasSentInRound returns true if the message tool sent a message during the current round.
|
|
||||||
func (t *MessageTool) HasSentInRound() bool {
|
|
||||||
return t.sentInRound
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *MessageTool) SetSendCallback(callback SendCallback) {
|
func (t *MessageTool) SetSendCallback(callback SendCallback) {
|
||||||
|
|
@ -71,6 +64,15 @@ func (t *MessageTool) Execute(ctx context.Context, args map[string]interface{})
|
||||||
channel, _ := args["channel"].(string)
|
channel, _ := args["channel"].(string)
|
||||||
chatID, _ := args["chat_id"].(string)
|
chatID, _ := args["chat_id"].(string)
|
||||||
|
|
||||||
|
if channel == "" || chatID == "" {
|
||||||
|
ctxChannel, ctxChatID := ExecutionTargetFromContext(ctx)
|
||||||
|
if channel == "" {
|
||||||
|
channel = ctxChannel
|
||||||
|
}
|
||||||
|
if chatID == "" {
|
||||||
|
chatID = ctxChatID
|
||||||
|
}
|
||||||
|
}
|
||||||
if channel == "" {
|
if channel == "" {
|
||||||
channel = t.defaultChannel
|
channel = t.defaultChannel
|
||||||
}
|
}
|
||||||
|
|
@ -94,7 +96,7 @@ func (t *MessageTool) Execute(ctx context.Context, args map[string]interface{})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
t.sentInRound = true
|
MarkMessageSent(ctx)
|
||||||
// Silent: user already received the message directly
|
// Silent: user already received the message directly
|
||||||
return &ToolResult{
|
return &ToolResult{
|
||||||
ForLLM: fmt.Sprintf("Message sent to %s:%s", channel, chatID),
|
ForLLM: fmt.Sprintf("Message sent to %s:%s", channel, chatID),
|
||||||
|
|
|
||||||
|
|
@ -79,19 +79,47 @@ func (t *ObligationTool) Parameters() map[string]interface{} {
|
||||||
},
|
},
|
||||||
"title": map[string]interface{}{
|
"title": map[string]interface{}{
|
||||||
"type": "string",
|
"type": "string",
|
||||||
"description": "Title for create action.",
|
"description": "Title for create action. Alias: content.",
|
||||||
|
},
|
||||||
|
"content": map[string]interface{}{
|
||||||
|
"type": "string",
|
||||||
|
"description": "Alias for title when creating an obligation.",
|
||||||
},
|
},
|
||||||
"details": map[string]interface{}{
|
"details": map[string]interface{}{
|
||||||
"type": "string",
|
"type": "string",
|
||||||
"description": "Optional details for create action.",
|
"description": "Optional details for create action. Aliases: notes, description.",
|
||||||
|
},
|
||||||
|
"notes": map[string]interface{}{
|
||||||
|
"type": "string",
|
||||||
|
"description": "Alias for details when creating an obligation.",
|
||||||
|
},
|
||||||
|
"description": map[string]interface{}{
|
||||||
|
"type": "string",
|
||||||
|
"description": "Alias for details when creating an obligation.",
|
||||||
},
|
},
|
||||||
"scheduled_at": map[string]interface{}{
|
"scheduled_at": map[string]interface{}{
|
||||||
"type": "string",
|
"type": "string",
|
||||||
"description": "Optional RFC3339 schedule time.",
|
"description": "Optional schedule time. Accepts RFC3339 and naive YYYY-MM-DDTHH:MM:SS. Aliases: remind_at, reminder_at.",
|
||||||
|
},
|
||||||
|
"remind_at": map[string]interface{}{
|
||||||
|
"type": "string",
|
||||||
|
"description": "Alias for scheduled_at when creating an obligation.",
|
||||||
|
},
|
||||||
|
"reminder_at": map[string]interface{}{
|
||||||
|
"type": "string",
|
||||||
|
"description": "Alias for scheduled_at when creating an obligation.",
|
||||||
},
|
},
|
||||||
"due_at": map[string]interface{}{
|
"due_at": map[string]interface{}{
|
||||||
"type": "string",
|
"type": "string",
|
||||||
"description": "Optional RFC3339 due time.",
|
"description": "Optional due time. Accepts RFC3339 and naive YYYY-MM-DDTHH:MM:SS. Aliases: due_date, deadline_at.",
|
||||||
|
},
|
||||||
|
"due_date": map[string]interface{}{
|
||||||
|
"type": "string",
|
||||||
|
"description": "Alias for due_at when creating an obligation.",
|
||||||
|
},
|
||||||
|
"deadline_at": map[string]interface{}{
|
||||||
|
"type": "string",
|
||||||
|
"description": "Alias for due_at when creating an obligation.",
|
||||||
},
|
},
|
||||||
"state": map[string]interface{}{
|
"state": map[string]interface{}{
|
||||||
"type": "string",
|
"type": "string",
|
||||||
|
|
@ -132,7 +160,7 @@ func (t *ObligationTool) Execute(ctx context.Context, args map[string]interface{
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *ObligationTool) create(ctx context.Context, args map[string]interface{}) *ToolResult {
|
func (t *ObligationTool) create(ctx context.Context, args map[string]interface{}) *ToolResult {
|
||||||
title, _ := args["title"].(string)
|
title := obligationFirstNonEmptyString(args, "title", "content")
|
||||||
if strings.TrimSpace(title) == "" {
|
if strings.TrimSpace(title) == "" {
|
||||||
return ErrorResult("title is required for create").WithError(fmt.Errorf("title is required"))
|
return ErrorResult("title is required for create").WithError(fmt.Errorf("title is required"))
|
||||||
}
|
}
|
||||||
|
|
@ -140,24 +168,24 @@ func (t *ObligationTool) create(ctx context.Context, args map[string]interface{}
|
||||||
rec := &ObligationRecord{
|
rec := &ObligationRecord{
|
||||||
ID: ids.New().String(),
|
ID: ids.New().String(),
|
||||||
Title: title,
|
Title: title,
|
||||||
Details: stringOr(args["details"]),
|
Details: obligationFirstNonEmptyString(args, "details", "notes", "description"),
|
||||||
State: ObligationStateCreated,
|
State: ObligationStateCreated,
|
||||||
CreatedAt: now,
|
CreatedAt: now,
|
||||||
UpdatedAt: now,
|
UpdatedAt: now,
|
||||||
}
|
}
|
||||||
|
|
||||||
if scheduledAtRaw := stringOr(args["scheduled_at"]); scheduledAtRaw != "" {
|
if scheduledAtRaw := obligationFirstNonEmptyString(args, "scheduled_at", "remind_at", "reminder_at"); scheduledAtRaw != "" {
|
||||||
ts, err := time.Parse(time.RFC3339, scheduledAtRaw)
|
ts, err := parseObligationTimestamp(scheduledAtRaw)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return ErrorResult("scheduled_at must be RFC3339").WithError(err)
|
return ErrorResult("scheduled_at must be RFC3339 or YYYY-MM-DDTHH:MM:SS").WithError(err)
|
||||||
}
|
}
|
||||||
rec.ScheduledAt = ts.UTC()
|
rec.ScheduledAt = ts.UTC()
|
||||||
rec.State = ObligationStateScheduled
|
rec.State = ObligationStateScheduled
|
||||||
}
|
}
|
||||||
if dueAtRaw := stringOr(args["due_at"]); dueAtRaw != "" {
|
if dueAtRaw := obligationFirstNonEmptyString(args, "due_at", "due_date", "deadline_at"); dueAtRaw != "" {
|
||||||
ts, err := time.Parse(time.RFC3339, dueAtRaw)
|
ts, err := parseObligationTimestamp(dueAtRaw)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return ErrorResult("due_at must be RFC3339").WithError(err)
|
return ErrorResult("due_at must be RFC3339 or YYYY-MM-DDTHH:MM:SS").WithError(err)
|
||||||
}
|
}
|
||||||
rec.DueAt = ts.UTC()
|
rec.DueAt = ts.UTC()
|
||||||
if rec.State == ObligationStateCreated {
|
if rec.State == ObligationStateCreated {
|
||||||
|
|
@ -437,3 +465,20 @@ func stringOr(v interface{}) string {
|
||||||
s, _ := v.(string)
|
s, _ := v.(string)
|
||||||
return s
|
return s
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func obligationFirstNonEmptyString(args map[string]interface{}, keys ...string) string {
|
||||||
|
for _, key := range keys {
|
||||||
|
if value := strings.TrimSpace(stringOr(args[key])); value != "" {
|
||||||
|
return value
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
func parseObligationTimestamp(raw string) (time.Time, error) {
|
||||||
|
ts, err := time.Parse(time.RFC3339, raw)
|
||||||
|
if err == nil {
|
||||||
|
return ts, nil
|
||||||
|
}
|
||||||
|
return time.Parse("2006-01-02T15:04:05", raw)
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -46,6 +46,56 @@ func TestObligationTool_CreateAndList(t *testing.T) {
|
||||||
assert.GreaterOrEqual(t, payload.Count, 1)
|
assert.GreaterOrEqual(t, payload.Count, 1)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestObligationTool_CreateAcceptsEvalStyleAliases(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
ctx := t.Context()
|
||||||
|
del, err := delegate.NewLibSQLInMemory()
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, del.Init(ctx))
|
||||||
|
defer del.Close()
|
||||||
|
|
||||||
|
tool := NewObligationTool(del, "test-agent")
|
||||||
|
create := tool.Execute(ctx, map[string]interface{}{
|
||||||
|
"action": "create",
|
||||||
|
"content": "Submit tax documents",
|
||||||
|
"details": "Critical financial deadline",
|
||||||
|
"remind_at": "2026-03-10T09:00:00",
|
||||||
|
"due_date": "2026-03-15T23:59:00",
|
||||||
|
"description": "backup description should be ignored when details present",
|
||||||
|
})
|
||||||
|
require.NotNil(t, create)
|
||||||
|
require.False(t, create.IsError, create.ForLLM)
|
||||||
|
|
||||||
|
var rec ObligationRecord
|
||||||
|
require.NoError(t, jsonv2.Unmarshal([]byte(create.ForLLM), &rec))
|
||||||
|
assert.Equal(t, "Submit tax documents", rec.Title)
|
||||||
|
assert.Equal(t, "Critical financial deadline", rec.Details)
|
||||||
|
assert.Empty(t, cmp.Diff(ObligationStateScheduled, rec.State))
|
||||||
|
assert.Equal(t, time.Date(2026, 3, 10, 9, 0, 0, 0, time.UTC), rec.ScheduledAt)
|
||||||
|
assert.Equal(t, time.Date(2026, 3, 15, 23, 59, 0, 0, time.UTC), rec.DueAt)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestObligationTool_ParametersDescribeCreateAliases(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
tool := NewObligationTool(nil, "test-agent")
|
||||||
|
params := tool.Parameters()
|
||||||
|
properties, ok := params["properties"].(map[string]interface{})
|
||||||
|
require.True(t, ok)
|
||||||
|
|
||||||
|
for _, key := range []string{"content", "notes", "description", "remind_at", "reminder_at", "due_date", "deadline_at"} {
|
||||||
|
_, ok := properties[key]
|
||||||
|
assert.True(t, ok, "expected parameters to include alias field %q", key)
|
||||||
|
}
|
||||||
|
|
||||||
|
titleDesc := properties["title"].(map[string]interface{})["description"].(string)
|
||||||
|
assert.Contains(t, titleDesc, "Alias: content")
|
||||||
|
dueDesc := properties["due_at"].(map[string]interface{})["description"].(string)
|
||||||
|
assert.Contains(t, dueDesc, "due_date")
|
||||||
|
scheduledDesc := properties["scheduled_at"].(map[string]interface{})["description"].(string)
|
||||||
|
assert.Contains(t, scheduledDesc, "remind_at")
|
||||||
|
}
|
||||||
|
|
||||||
func TestObligationTool_StateMachineAndEvidence(t *testing.T) {
|
func TestObligationTool_StateMachineAndEvidence(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
ctx := t.Context()
|
ctx := t.Context()
|
||||||
|
|
|
||||||
|
|
@ -83,15 +83,17 @@ func (r *ToolRegistry) Execute(ctx context.Context, name string, args map[string
|
||||||
return r.ExecuteWithContext(ctx, name, args, "", "", nil)
|
return r.ExecuteWithContext(ctx, name, args, "", "", nil)
|
||||||
}
|
}
|
||||||
|
|
||||||
// ExecuteWithContext executes a tool with channel/chatID context and optional async callback.
|
// ExecuteWithContext executes a tool with channel/chatID context and optional
|
||||||
// If the tool implements AsyncTool and a non-nil callback is provided,
|
// async callback. Per-call execution metadata is carried in the context so
|
||||||
// the callback will be set on the tool before execution.
|
// singleton tool instances are not mutated on the hot path.
|
||||||
func (r *ToolRegistry) ExecuteWithContext(ctx context.Context, name string, args map[string]interface{}, channel, chatID string, asyncCallback AsyncCallback) *ToolResult {
|
func (r *ToolRegistry) ExecuteWithContext(ctx context.Context, name string, args map[string]interface{}, channel, chatID string, asyncCallback AsyncCallback) *ToolResult {
|
||||||
logger.InfoCF("tool", "Tool execution started",
|
logger.InfoCF("tool", "Tool execution started",
|
||||||
map[string]interface{}{
|
map[string]interface{}{
|
||||||
"tool": name,
|
"tool": name,
|
||||||
"args": args,
|
"args": args,
|
||||||
})
|
})
|
||||||
|
ctx = WithExecutionTarget(ctx, channel, chatID)
|
||||||
|
ctx = WithAsyncCallback(ctx, asyncCallback)
|
||||||
|
|
||||||
tool, ok := r.Get(name)
|
tool, ok := r.Get(name)
|
||||||
if !ok {
|
if !ok {
|
||||||
|
|
@ -102,18 +104,13 @@ func (r *ToolRegistry) ExecuteWithContext(ctx context.Context, name string, args
|
||||||
return ErrorResult(fmt.Sprintf("tool %q not found", name)).WithError(fmt.Errorf("tool not found"))
|
return ErrorResult(fmt.Sprintf("tool %q not found", name)).WithError(fmt.Errorf("tool not found"))
|
||||||
}
|
}
|
||||||
|
|
||||||
// If tool implements ContextualTool, set context
|
// Backward-compatible bridge for tools that still implement the legacy hook
|
||||||
|
// interfaces instead of reading execution metadata from context directly.
|
||||||
if contextualTool, ok := tool.(ContextualTool); ok && channel != "" && chatID != "" {
|
if contextualTool, ok := tool.(ContextualTool); ok && channel != "" && chatID != "" {
|
||||||
contextualTool.SetContext(channel, chatID)
|
contextualTool.SetContext(channel, chatID)
|
||||||
}
|
}
|
||||||
|
if asyncTool, ok := tool.(AsyncTool); ok {
|
||||||
// If tool implements AsyncTool and callback is provided, set callback
|
|
||||||
if asyncTool, ok := tool.(AsyncTool); ok && asyncCallback != nil {
|
|
||||||
asyncTool.SetCallback(asyncCallback)
|
asyncTool.SetCallback(asyncCallback)
|
||||||
logger.DebugCF("tool", "Async callback injected",
|
|
||||||
map[string]interface{}{
|
|
||||||
"tool": name,
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
start := time.Now()
|
start := time.Now()
|
||||||
|
|
|
||||||
|
|
@ -149,7 +149,7 @@ func (t *ToolSearchTool) focusTerms(ctx context.Context) []string {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
sessionKey := t.sessionKeyFn()
|
sessionKey := ResolveSessionKey(ctx, t.sessionKeyFn)
|
||||||
if strings.TrimSpace(sessionKey) == "" {
|
if strings.TrimSpace(sessionKey) == "" {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
|
||||||
134
pkg/tools/session_context.go
Normal file
134
pkg/tools/session_context.go
Normal file
|
|
@ -0,0 +1,134 @@
|
||||||
|
package tools
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"strings"
|
||||||
|
"sync/atomic"
|
||||||
|
)
|
||||||
|
|
||||||
|
type ctxSessionKey struct{}
|
||||||
|
type ctxExecutionTarget struct{}
|
||||||
|
type ctxAsyncCallback struct{}
|
||||||
|
type ctxMessageSendTracker struct{}
|
||||||
|
|
||||||
|
type executionTarget struct {
|
||||||
|
channel string
|
||||||
|
chatID string
|
||||||
|
}
|
||||||
|
|
||||||
|
type MessageSendTracker struct {
|
||||||
|
sent atomic.Bool
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithSessionKey annotates the execution context with the session key that
|
||||||
|
// session-scoped tools should use for this call tree.
|
||||||
|
func WithSessionKey(ctx context.Context, sessionKey string) context.Context {
|
||||||
|
if strings.TrimSpace(sessionKey) == "" {
|
||||||
|
return ctx
|
||||||
|
}
|
||||||
|
return context.WithValue(ctx, ctxSessionKey{}, sessionKey)
|
||||||
|
}
|
||||||
|
|
||||||
|
// SessionKeyFromContext returns the session key previously attached via
|
||||||
|
// WithSessionKey. Missing values are treated as empty.
|
||||||
|
func SessionKeyFromContext(ctx context.Context) string {
|
||||||
|
v, _ := ctx.Value(ctxSessionKey{}).(string)
|
||||||
|
return strings.TrimSpace(v)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ResolveSessionKey prefers an explicit session key carried in context and
|
||||||
|
// falls back to a legacy resolver when no per-call session key is present.
|
||||||
|
func ResolveSessionKey(ctx context.Context, fallback func() string) string {
|
||||||
|
if sessionKey := SessionKeyFromContext(ctx); sessionKey != "" {
|
||||||
|
return sessionKey
|
||||||
|
}
|
||||||
|
if fallback != nil {
|
||||||
|
return strings.TrimSpace(fallback())
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithExecutionTarget annotates the execution context with the channel/chat
|
||||||
|
// destination that contextual and async tools should treat as the current user
|
||||||
|
// target.
|
||||||
|
func WithExecutionTarget(ctx context.Context, channel, chatID string) context.Context {
|
||||||
|
channel = strings.TrimSpace(channel)
|
||||||
|
chatID = strings.TrimSpace(chatID)
|
||||||
|
if channel == "" && chatID == "" {
|
||||||
|
return ctx
|
||||||
|
}
|
||||||
|
return context.WithValue(ctx, ctxExecutionTarget{}, executionTarget{channel: channel, chatID: chatID})
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExecutionTargetFromContext returns the channel/chat destination previously
|
||||||
|
// attached via WithExecutionTarget.
|
||||||
|
func ExecutionTargetFromContext(ctx context.Context) (channel, chatID string) {
|
||||||
|
v, _ := ctx.Value(ctxExecutionTarget{}).(executionTarget)
|
||||||
|
return v.channel, v.chatID
|
||||||
|
}
|
||||||
|
|
||||||
|
// ResolveExecutionTarget prefers a target carried in context and falls back to
|
||||||
|
// the provided defaults when the context does not specify one.
|
||||||
|
func ResolveExecutionTarget(ctx context.Context, fallbackChannel, fallbackChatID string) (string, string) {
|
||||||
|
channel, chatID := ExecutionTargetFromContext(ctx)
|
||||||
|
if channel == "" {
|
||||||
|
channel = strings.TrimSpace(fallbackChannel)
|
||||||
|
}
|
||||||
|
if chatID == "" {
|
||||||
|
chatID = strings.TrimSpace(fallbackChatID)
|
||||||
|
}
|
||||||
|
return channel, chatID
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithAsyncCallback annotates the execution context with the async completion
|
||||||
|
// callback that async tools should invoke for background completions.
|
||||||
|
func WithAsyncCallback(ctx context.Context, cb AsyncCallback) context.Context {
|
||||||
|
if cb == nil {
|
||||||
|
return ctx
|
||||||
|
}
|
||||||
|
return context.WithValue(ctx, ctxAsyncCallback{}, cb)
|
||||||
|
}
|
||||||
|
|
||||||
|
// AsyncCallbackFromContext returns the async completion callback previously
|
||||||
|
// attached via WithAsyncCallback.
|
||||||
|
func AsyncCallbackFromContext(ctx context.Context) AsyncCallback {
|
||||||
|
v, _ := ctx.Value(ctxAsyncCallback{}).(AsyncCallback)
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewMessageSendTracker allocates per-execution message-send state used to
|
||||||
|
// suppress duplicate final replies when a direct message already went out.
|
||||||
|
func NewMessageSendTracker() *MessageSendTracker {
|
||||||
|
return &MessageSendTracker{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithMessageSendTracker attaches a per-execution message-send tracker.
|
||||||
|
func WithMessageSendTracker(ctx context.Context, tracker *MessageSendTracker) context.Context {
|
||||||
|
if tracker == nil {
|
||||||
|
return ctx
|
||||||
|
}
|
||||||
|
return context.WithValue(ctx, ctxMessageSendTracker{}, tracker)
|
||||||
|
}
|
||||||
|
|
||||||
|
// MessageSendTrackerFromContext returns the tracker attached via
|
||||||
|
// WithMessageSendTracker, if present.
|
||||||
|
func MessageSendTrackerFromContext(ctx context.Context) *MessageSendTracker {
|
||||||
|
v, _ := ctx.Value(ctxMessageSendTracker{}).(*MessageSendTracker)
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
// MarkMessageSent records that the message tool already sent a direct reply in
|
||||||
|
// the current execution.
|
||||||
|
func MarkMessageSent(ctx context.Context) {
|
||||||
|
if tracker := MessageSendTrackerFromContext(ctx); tracker != nil {
|
||||||
|
tracker.sent.Store(true)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Sent reports whether a direct message was already sent in the tracked round.
|
||||||
|
func (t *MessageSendTracker) Sent() bool {
|
||||||
|
if t == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return t.sent.Load()
|
||||||
|
}
|
||||||
|
|
@ -76,9 +76,14 @@ func (t *SpawnTool) Execute(ctx context.Context, args map[string]interface{}) *T
|
||||||
if t.manager == nil {
|
if t.manager == nil {
|
||||||
return ErrorResult("Subagent manager not configured")
|
return ErrorResult("Subagent manager not configured")
|
||||||
}
|
}
|
||||||
|
originChannel, originChatID := ResolveExecutionTarget(ctx, t.originChannel, t.originChatID)
|
||||||
|
callback := AsyncCallbackFromContext(ctx)
|
||||||
|
if callback == nil && (originChannel == "" || originChatID == "") {
|
||||||
|
callback = t.callback
|
||||||
|
}
|
||||||
|
|
||||||
// Pass callback to manager for async completion notification
|
// Pass callback to manager for async completion notification
|
||||||
result, err := t.manager.Spawn(ctx, task, label, delegatedScope, keptWork, t.originChannel, t.originChatID, t.callback)
|
result, err := t.manager.Spawn(ctx, task, label, delegatedScope, keptWork, originChannel, originChatID, callback)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return ErrorResult(fmt.Sprintf("failed to spawn subagent: %v", err))
|
return ErrorResult(fmt.Sprintf("failed to spawn subagent: %v", err))
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -34,8 +34,9 @@ type RunLoopFunc func(ctx context.Context, config ToolLoopConfig, systemPrompt,
|
||||||
type delegationCtxKey string
|
type delegationCtxKey string
|
||||||
|
|
||||||
const (
|
const (
|
||||||
delegationTaskIDKey delegationCtxKey = "delegation_task_id"
|
delegationTaskIDKey delegationCtxKey = "delegation_task_id"
|
||||||
delegationDepthKey delegationCtxKey = "delegation_depth"
|
delegationDepthKey delegationCtxKey = "delegation_depth"
|
||||||
|
delegationSessionKey delegationCtxKey = "delegation_session_key"
|
||||||
)
|
)
|
||||||
|
|
||||||
func delegationTaskIDFromContext(ctx context.Context) string {
|
func delegationTaskIDFromContext(ctx context.Context) string {
|
||||||
|
|
@ -52,12 +53,30 @@ func delegationDepthFromContext(ctx context.Context) int {
|
||||||
return 0
|
return 0
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func delegationSessionKeyFromContext(ctx context.Context) string {
|
||||||
|
if v, ok := ctx.Value(delegationSessionKey).(string); ok {
|
||||||
|
return strings.TrimSpace(v)
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
func DelegationSessionKeyFromContext(ctx context.Context) string {
|
||||||
|
return delegationSessionKeyFromContext(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
func withDelegationContext(ctx context.Context, taskID string, depth int) context.Context {
|
func withDelegationContext(ctx context.Context, taskID string, depth int) context.Context {
|
||||||
ctx = context.WithValue(ctx, delegationTaskIDKey, taskID)
|
ctx = context.WithValue(ctx, delegationTaskIDKey, taskID)
|
||||||
ctx = context.WithValue(ctx, delegationDepthKey, depth)
|
ctx = context.WithValue(ctx, delegationDepthKey, depth)
|
||||||
return ctx
|
return ctx
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func withDelegationSessionKey(ctx context.Context, sessionKey string) context.Context {
|
||||||
|
if strings.TrimSpace(sessionKey) == "" {
|
||||||
|
return ctx
|
||||||
|
}
|
||||||
|
return context.WithValue(ctx, delegationSessionKey, sessionKey)
|
||||||
|
}
|
||||||
|
|
||||||
// DelegationAuditEvent captures lineage and outcomes for delegated work.
|
// DelegationAuditEvent captures lineage and outcomes for delegated work.
|
||||||
type DelegationAuditEvent struct {
|
type DelegationAuditEvent struct {
|
||||||
TaskID string
|
TaskID string
|
||||||
|
|
@ -233,8 +252,9 @@ func (sm *SubagentManager) Spawn(ctx context.Context, task, label, delegatedScop
|
||||||
OriginChatID: originChatID,
|
OriginChatID: originChatID,
|
||||||
})
|
})
|
||||||
|
|
||||||
// Start task in background with context cancellation support
|
// Start task in background with context cancellation support.
|
||||||
go sm.runTask(ctx, subagentTask, callback)
|
taskCtx := withDelegationSessionKey(ctx, SessionKeyFromContext(ctx))
|
||||||
|
go sm.runTask(taskCtx, subagentTask, callback)
|
||||||
|
|
||||||
if label != "" {
|
if label != "" {
|
||||||
return fmt.Sprintf("Spawned subagent '%s' for task: %s", label, task), nil
|
return fmt.Sprintf("Spawned subagent '%s' for task: %s", label, task), nil
|
||||||
|
|
@ -271,7 +291,7 @@ Complete the task independently and provide a clear summary of what was done.`
|
||||||
if runLoop == nil {
|
if runLoop == nil {
|
||||||
err = ErrRunLoopNotConfigured
|
err = ErrRunLoopNotConfigured
|
||||||
} else {
|
} else {
|
||||||
taskCtx := withDelegationContext(ctx, task.ID, task.Depth)
|
taskCtx := withDelegationSessionKey(withDelegationContext(ctx, task.ID, task.Depth), delegationSessionKeyFromContext(ctx))
|
||||||
loopResult, err = runLoop(taskCtx, ToolLoopConfig{
|
loopResult, err = runLoop(taskCtx, ToolLoopConfig{
|
||||||
Model: sm.model,
|
Model: sm.model,
|
||||||
ModelID: sm.defaultModel,
|
ModelID: sm.defaultModel,
|
||||||
|
|
@ -449,6 +469,7 @@ func (t *SubagentTool) Execute(ctx context.Context, args map[string]interface{})
|
||||||
if t.manager == nil {
|
if t.manager == nil {
|
||||||
return ErrorResult("Subagent manager not configured").WithError(fmt.Errorf("manager is nil"))
|
return ErrorResult("Subagent manager not configured").WithError(fmt.Errorf("manager is nil"))
|
||||||
}
|
}
|
||||||
|
originChannel, originChatID := ResolveExecutionTarget(ctx, t.originChannel, t.originChatID)
|
||||||
|
|
||||||
sm := t.manager
|
sm := t.manager
|
||||||
parentTaskID := delegationTaskIDFromContext(ctx)
|
parentTaskID := delegationTaskIDFromContext(ctx)
|
||||||
|
|
@ -499,7 +520,7 @@ func (t *SubagentTool) Execute(ctx context.Context, args map[string]interface{})
|
||||||
}()
|
}()
|
||||||
|
|
||||||
taskID := fmt.Sprintf("subagent-sync-%d", time.Now().UnixNano())
|
taskID := fmt.Sprintf("subagent-sync-%d", time.Now().UnixNano())
|
||||||
taskCtx := withDelegationContext(ctx, taskID, childDepth)
|
taskCtx := withDelegationSessionKey(withDelegationContext(ctx, taskID, childDepth), SessionKeyFromContext(ctx))
|
||||||
sm.emitAudit(ctx, DelegationAuditEvent{
|
sm.emitAudit(ctx, DelegationAuditEvent{
|
||||||
TaskID: taskID,
|
TaskID: taskID,
|
||||||
ParentTaskID: parentTaskID,
|
ParentTaskID: parentTaskID,
|
||||||
|
|
@ -509,8 +530,8 @@ func (t *SubagentTool) Execute(ctx context.Context, args map[string]interface{})
|
||||||
Label: label,
|
Label: label,
|
||||||
DelegatedScope: delegatedScope,
|
DelegatedScope: delegatedScope,
|
||||||
KeptWork: keptWork,
|
KeptWork: keptWork,
|
||||||
OriginChannel: t.originChannel,
|
OriginChannel: originChannel,
|
||||||
OriginChatID: t.originChatID,
|
OriginChatID: originChatID,
|
||||||
})
|
})
|
||||||
|
|
||||||
runLoop := sm.getRunLoop()
|
runLoop := sm.getRunLoop()
|
||||||
|
|
@ -523,7 +544,7 @@ func (t *SubagentTool) Execute(ctx context.Context, args map[string]interface{})
|
||||||
Tools: tools,
|
Tools: tools,
|
||||||
Bus: sm.bus,
|
Bus: sm.bus,
|
||||||
MaxIterations: maxIter,
|
MaxIterations: maxIter,
|
||||||
}, systemPrompt, task, t.originChannel, t.originChatID)
|
}, systemPrompt, task, originChannel, originChatID)
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
sm.emitAudit(ctx, DelegationAuditEvent{
|
sm.emitAudit(ctx, DelegationAuditEvent{
|
||||||
|
|
@ -536,8 +557,8 @@ func (t *SubagentTool) Execute(ctx context.Context, args map[string]interface{})
|
||||||
DelegatedScope: delegatedScope,
|
DelegatedScope: delegatedScope,
|
||||||
KeptWork: keptWork,
|
KeptWork: keptWork,
|
||||||
Error: err.Error(),
|
Error: err.Error(),
|
||||||
OriginChannel: t.originChannel,
|
OriginChannel: originChannel,
|
||||||
OriginChatID: t.originChatID,
|
OriginChatID: originChatID,
|
||||||
})
|
})
|
||||||
return ErrorResult(fmt.Sprintf("Subagent execution failed: %v", err)).WithError(err)
|
return ErrorResult(fmt.Sprintf("Subagent execution failed: %v", err)).WithError(err)
|
||||||
}
|
}
|
||||||
|
|
@ -567,8 +588,8 @@ func (t *SubagentTool) Execute(ctx context.Context, args map[string]interface{})
|
||||||
KeptWork: keptWork,
|
KeptWork: keptWork,
|
||||||
Iterations: loopResult.Iterations,
|
Iterations: loopResult.Iterations,
|
||||||
ResultChars: len(loopResult.Content),
|
ResultChars: len(loopResult.Content),
|
||||||
OriginChannel: t.originChannel,
|
OriginChannel: originChannel,
|
||||||
OriginChatID: t.originChatID,
|
OriginChatID: originChatID,
|
||||||
})
|
})
|
||||||
|
|
||||||
return &ToolResult{
|
return &ToolResult{
|
||||||
|
|
|
||||||
|
|
@ -18,6 +18,8 @@ const (
|
||||||
userAgent = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36"
|
userAgent = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
var htmlTitleRE = regexp.MustCompile(`(?is)<title[^>]*>(.*?)</title>`)
|
||||||
|
|
||||||
type SearchProvider interface {
|
type SearchProvider interface {
|
||||||
Search(ctx context.Context, query string, count int) (string, error)
|
Search(ctx context.Context, query string, count int) (string, error)
|
||||||
}
|
}
|
||||||
|
|
@ -467,6 +469,9 @@ func (t *WebFetchTool) Execute(ctx context.Context, args map[string]interface{})
|
||||||
|
|
||||||
var text, extractor string
|
var text, extractor string
|
||||||
|
|
||||||
|
isHTML := strings.Contains(contentType, "text/html") || len(body) > 0 &&
|
||||||
|
(strings.HasPrefix(string(body), "<!DOCTYPE") || strings.HasPrefix(strings.ToLower(string(body)), "<html"))
|
||||||
|
|
||||||
if strings.Contains(contentType, "application/json") {
|
if strings.Contains(contentType, "application/json") {
|
||||||
var jsonData interface{}
|
var jsonData interface{}
|
||||||
if err := jsonv2.Unmarshal(body, &jsonData); err == nil {
|
if err := jsonv2.Unmarshal(body, &jsonData); err == nil {
|
||||||
|
|
@ -477,8 +482,7 @@ func (t *WebFetchTool) Execute(ctx context.Context, args map[string]interface{})
|
||||||
text = string(body)
|
text = string(body)
|
||||||
extractor = "raw"
|
extractor = "raw"
|
||||||
}
|
}
|
||||||
} else if strings.Contains(contentType, "text/html") || len(body) > 0 &&
|
} else if isHTML {
|
||||||
(strings.HasPrefix(string(body), "<!DOCTYPE") || strings.HasPrefix(strings.ToLower(string(body)), "<html")) {
|
|
||||||
text = t.extractText(string(body))
|
text = t.extractText(string(body))
|
||||||
extractor = "text"
|
extractor = "text"
|
||||||
} else {
|
} else {
|
||||||
|
|
@ -501,20 +505,47 @@ func (t *WebFetchTool) Execute(ctx context.Context, args map[string]interface{})
|
||||||
}
|
}
|
||||||
|
|
||||||
resultJSON, _ := jsonv2.Marshal(result, jsontext.WithIndent(" "))
|
resultJSON, _ := jsonv2.Marshal(result, jsontext.WithIndent(" "))
|
||||||
|
llmContent := formatWebFetchLLMContent(urlStr, resp.StatusCode, extractor, truncated, text, isHTML, string(body))
|
||||||
|
|
||||||
return &ToolResult{
|
return &ToolResult{
|
||||||
ForLLM: fmt.Sprintf("Fetched %d bytes from %s (extractor: %s, truncated: %v)", len(text), urlStr, extractor, truncated),
|
ForLLM: llmContent,
|
||||||
ForUser: string(resultJSON),
|
ForUser: string(resultJSON),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func formatWebFetchLLMContent(urlStr string, status int, extractor string, truncated bool, text string, isHTML bool, rawHTML string) string {
|
||||||
|
sections := []string{fmt.Sprintf("Fetched %s (status %d, extractor: %s, truncated: %v)", urlStr, status, extractor, truncated)}
|
||||||
|
|
||||||
|
if isHTML {
|
||||||
|
if title := extractHTMLTitle(rawHTML); title != "" {
|
||||||
|
sections = append(sections, "Title: "+title)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if trimmed := strings.TrimSpace(text); trimmed != "" {
|
||||||
|
sections = append(sections, "Content:\n"+trimmed)
|
||||||
|
}
|
||||||
|
|
||||||
|
return strings.Join(sections, "\n\n")
|
||||||
|
}
|
||||||
|
|
||||||
|
func extractHTMLTitle(htmlContent string) string {
|
||||||
|
m := htmlTitleRE.FindStringSubmatch(htmlContent)
|
||||||
|
if len(m) < 2 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
title := stripTags(m[1])
|
||||||
|
title = strings.Join(strings.Fields(title), " ")
|
||||||
|
return strings.TrimSpace(title)
|
||||||
|
}
|
||||||
|
|
||||||
func (t *WebFetchTool) extractText(htmlContent string) string {
|
func (t *WebFetchTool) extractText(htmlContent string) string {
|
||||||
re := regexp.MustCompile(`<script[\s\S]*?</script>`)
|
re := regexp.MustCompile(`<script[\s\S]*?</script>`)
|
||||||
result := re.ReplaceAllLiteralString(htmlContent, "")
|
result := re.ReplaceAllLiteralString(htmlContent, "")
|
||||||
re = regexp.MustCompile(`<style[\s\S]*?</style>`)
|
re = regexp.MustCompile(`<style[\s\S]*?</style>`)
|
||||||
result = re.ReplaceAllLiteralString(result, "")
|
result = re.ReplaceAllLiteralString(result, "")
|
||||||
re = regexp.MustCompile(`<[^>]+>`)
|
re = regexp.MustCompile(`<[^>]+>`)
|
||||||
result = re.ReplaceAllLiteralString(result, "")
|
result = re.ReplaceAllLiteralString(result, " ")
|
||||||
|
|
||||||
result = strings.TrimSpace(result)
|
result = strings.TrimSpace(result)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -16,7 +16,7 @@ func TestWebTool_WebFetch_Success(t *testing.T) {
|
||||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
w.Header().Set("Content-Type", "text/html")
|
w.Header().Set("Content-Type", "text/html")
|
||||||
w.WriteHeader(http.StatusOK)
|
w.WriteHeader(http.StatusOK)
|
||||||
w.Write([]byte("<html><body><h1>Test Page</h1><p>Content here</p></body></html>"))
|
w.Write([]byte("<html><head><title>Test Page</title></head><body><h1>Test Page</h1><p>Content here</p></body></html>"))
|
||||||
}))
|
}))
|
||||||
defer server.Close()
|
defer server.Close()
|
||||||
|
|
||||||
|
|
@ -42,6 +42,14 @@ func TestWebTool_WebFetch_Success(t *testing.T) {
|
||||||
if !strings.Contains(result.ForLLM, "bytes") && !strings.Contains(result.ForLLM, "extractor") {
|
if !strings.Contains(result.ForLLM, "bytes") && !strings.Contains(result.ForLLM, "extractor") {
|
||||||
t.Errorf("Expected ForLLM to contain summary, got: %s", result.ForLLM)
|
t.Errorf("Expected ForLLM to contain summary, got: %s", result.ForLLM)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if !strings.Contains(result.ForLLM, "Title: Test Page") {
|
||||||
|
t.Errorf("Expected ForLLM to contain extracted title, got: %s", result.ForLLM)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !strings.Contains(result.ForLLM, "Content here") {
|
||||||
|
t.Errorf("Expected ForLLM to contain fetched content, got: %s", result.ForLLM)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestWebTool_WebFetch_JSON verifies JSON content handling
|
// TestWebTool_WebFetch_JSON verifies JSON content handling
|
||||||
|
|
@ -216,7 +224,7 @@ func TestWebTool_WebFetch_HTMLExtraction(t *testing.T) {
|
||||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
w.Header().Set("Content-Type", "text/html")
|
w.Header().Set("Content-Type", "text/html")
|
||||||
w.WriteHeader(http.StatusOK)
|
w.WriteHeader(http.StatusOK)
|
||||||
w.Write([]byte(`<html><body><script>alert('test');</script><style>body{color:red;}</style><h1>Title</h1><p>Content</p></body></html>`))
|
w.Write([]byte(`<html><head><title>Title</title></head><body><script>alert('test');</script><style>body{color:red;}</style><h1>Title</h1><p>Content</p></body></html>`))
|
||||||
}))
|
}))
|
||||||
defer server.Close()
|
defer server.Close()
|
||||||
|
|
||||||
|
|
@ -238,6 +246,10 @@ func TestWebTool_WebFetch_HTMLExtraction(t *testing.T) {
|
||||||
t.Errorf("Expected ForUser to contain extracted text, got: %s", result.ForUser)
|
t.Errorf("Expected ForUser to contain extracted text, got: %s", result.ForUser)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if !strings.Contains(result.ForLLM, "Title: Title") {
|
||||||
|
t.Errorf("Expected ForLLM to contain extracted title, got: %s", result.ForLLM)
|
||||||
|
}
|
||||||
|
|
||||||
// Should NOT contain script or style tags
|
// Should NOT contain script or style tags
|
||||||
if strings.Contains(result.ForUser, "<script>") || strings.Contains(result.ForUser, "<style>") {
|
if strings.Contains(result.ForUser, "<script>") || strings.Contains(result.ForUser, "<style>") {
|
||||||
t.Errorf("Expected script/style tags to be removed, got: %s", result.ForUser)
|
t.Errorf("Expected script/style tags to be removed, got: %s", result.ForUser)
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue