updates on the ai and agent and all thse things awith onboarding
This commit is contained in:
238
internal/ai/playground/openai_compat.go
Normal file
238
internal/ai/playground/openai_compat.go
Normal file
@@ -0,0 +1,238 @@
|
||||
package playground
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// OpenAICompat drives the playground through any OpenAI-compatible chat
|
||||
// completions API — Groq by default (https://api.groq.com/openai/v1), or xAI,
|
||||
// or another provider — with plain net/http, so no SDK dependency is needed.
|
||||
//
|
||||
// Configured from PLAYGROUND_LLM_BASE_URL / _API_KEY / _MODEL (see config).
|
||||
// The model comes from that setting, not from the agent's registry pin: the
|
||||
// registry holds Claude ids for AI_engine, which this provider cannot serve.
|
||||
type OpenAICompat struct {
|
||||
BaseURL string
|
||||
APIKey string
|
||||
Model string
|
||||
HTTP *http.Client
|
||||
}
|
||||
|
||||
// NewOpenAICompat returns a client with a bounded HTTP timeout.
|
||||
func NewOpenAICompat(baseURL, apiKey, model string) *OpenAICompat {
|
||||
return &OpenAICompat{
|
||||
BaseURL: strings.TrimRight(baseURL, "/"),
|
||||
APIKey: apiKey,
|
||||
Model: model,
|
||||
HTTP: &http.Client{Timeout: 90 * time.Second},
|
||||
}
|
||||
}
|
||||
|
||||
// ModelName is what the trace reports as the model that answered.
|
||||
func (o *OpenAICompat) ModelName() string { return o.Model }
|
||||
|
||||
// ── Wire types (OpenAI chat completions) ────────────────────────────────────
|
||||
|
||||
type oaFunctionCall struct {
|
||||
Name string `json:"name"`
|
||||
Arguments string `json:"arguments"`
|
||||
}
|
||||
|
||||
type oaToolCall struct {
|
||||
ID string `json:"id"`
|
||||
Type string `json:"type"`
|
||||
Function oaFunctionCall `json:"function"`
|
||||
}
|
||||
|
||||
type oaMessage struct {
|
||||
Role string `json:"role"`
|
||||
Content *string `json:"content"`
|
||||
ToolCalls []oaToolCall `json:"tool_calls,omitempty"`
|
||||
ToolCallID string `json:"tool_call_id,omitempty"`
|
||||
}
|
||||
|
||||
type oaTool struct {
|
||||
Type string `json:"type"`
|
||||
Function struct {
|
||||
Name string `json:"name"`
|
||||
Description string `json:"description"`
|
||||
Parameters json.RawMessage `json:"parameters"`
|
||||
} `json:"function"`
|
||||
}
|
||||
|
||||
type oaRequest struct {
|
||||
Model string `json:"model"`
|
||||
Messages []oaMessage `json:"messages"`
|
||||
Tools []oaTool `json:"tools,omitempty"`
|
||||
MaxCompletionTokens int64 `json:"max_completion_tokens,omitempty"`
|
||||
}
|
||||
|
||||
type oaResponse struct {
|
||||
Choices []struct {
|
||||
Message oaMessage `json:"message"`
|
||||
FinishReason string `json:"finish_reason"`
|
||||
} `json:"choices"`
|
||||
Usage struct {
|
||||
PromptTokens int64 `json:"prompt_tokens"`
|
||||
CompletionTokens int64 `json:"completion_tokens"`
|
||||
} `json:"usage"`
|
||||
Error *struct {
|
||||
Message string `json:"message"`
|
||||
} `json:"error"`
|
||||
}
|
||||
|
||||
func strp(s string) *string { return &s }
|
||||
|
||||
// objectSchema makes sure a tool's parameters are an object schema with a
|
||||
// properties map — OpenAI-style APIs reject `{"type":"object"}` alone on some
|
||||
// models, and several registry tools have open schemas.
|
||||
func objectSchema(raw json.RawMessage) json.RawMessage {
|
||||
var m map[string]any
|
||||
if len(raw) == 0 || json.Unmarshal(raw, &m) != nil || m == nil {
|
||||
m = map[string]any{}
|
||||
}
|
||||
m["type"] = "object"
|
||||
if _, ok := m["properties"]; !ok {
|
||||
m["properties"] = map[string]any{}
|
||||
}
|
||||
b, _ := json.Marshal(m)
|
||||
return b
|
||||
}
|
||||
|
||||
// toWire converts the playground conversation into chat-completions messages.
|
||||
func toWire(req Request) oaRequest {
|
||||
out := oaRequest{MaxCompletionTokens: req.MaxTokens}
|
||||
if req.System != "" {
|
||||
out.Messages = append(out.Messages, oaMessage{Role: "system", Content: strp(req.System)})
|
||||
}
|
||||
for _, t := range req.Turns {
|
||||
switch {
|
||||
case t.Role == "assistant":
|
||||
msg := oaMessage{Role: "assistant"}
|
||||
var texts []string
|
||||
for _, b := range t.Assistant {
|
||||
switch {
|
||||
case b.Type == "text" && b.Text != "":
|
||||
texts = append(texts, b.Text)
|
||||
case b.Type == "tool_use" && b.ToolUse != nil:
|
||||
args := string(b.ToolUse.Input)
|
||||
if args == "" {
|
||||
args = "{}"
|
||||
}
|
||||
msg.ToolCalls = append(msg.ToolCalls, oaToolCall{
|
||||
ID: b.ToolUse.ID, Type: "function",
|
||||
Function: oaFunctionCall{Name: b.ToolUse.Name, Arguments: args},
|
||||
})
|
||||
}
|
||||
}
|
||||
if len(texts) > 0 {
|
||||
msg.Content = strp(strings.Join(texts, "\n\n"))
|
||||
}
|
||||
out.Messages = append(out.Messages, msg)
|
||||
case len(t.Results) > 0:
|
||||
for _, r := range t.Results {
|
||||
out.Messages = append(out.Messages, oaMessage{Role: "tool", ToolCallID: r.ToolUseID, Content: strp(r.Content)})
|
||||
}
|
||||
default:
|
||||
out.Messages = append(out.Messages, oaMessage{Role: "user", Content: strp(t.Text)})
|
||||
}
|
||||
}
|
||||
for _, td := range req.Tools {
|
||||
var tool oaTool
|
||||
tool.Type = "function"
|
||||
tool.Function.Name = td.Name
|
||||
tool.Function.Description = td.Description
|
||||
tool.Function.Parameters = objectSchema(td.InputSchema)
|
||||
out.Tools = append(out.Tools, tool)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// fromWire converts one chat-completions response into a playground Reply.
|
||||
func fromWire(resp oaResponse) (Reply, error) {
|
||||
if len(resp.Choices) == 0 {
|
||||
return Reply{}, errors.New("model returned no choices")
|
||||
}
|
||||
ch := resp.Choices[0]
|
||||
r := Reply{InputTokens: resp.Usage.PromptTokens, OutputTokens: resp.Usage.CompletionTokens}
|
||||
if ch.Message.Content != nil && strings.TrimSpace(*ch.Message.Content) != "" {
|
||||
r.Blocks = append(r.Blocks, Block{Type: "text", Text: *ch.Message.Content})
|
||||
}
|
||||
for _, tc := range ch.Message.ToolCalls {
|
||||
input := json.RawMessage(tc.Function.Arguments)
|
||||
if !json.Valid(input) {
|
||||
input = json.RawMessage(`{}`)
|
||||
}
|
||||
r.Blocks = append(r.Blocks, Block{Type: "tool_use", ToolUse: &ToolUse{ID: tc.ID, Name: tc.Function.Name, Input: input}})
|
||||
}
|
||||
switch ch.FinishReason {
|
||||
case "tool_calls":
|
||||
r.StopReason = "tool_use"
|
||||
case "length":
|
||||
r.StopReason = "max_tokens"
|
||||
default:
|
||||
r.StopReason = "end_turn"
|
||||
}
|
||||
// Some providers report "stop" while still returning tool calls; the calls win.
|
||||
if len(ch.Message.ToolCalls) > 0 {
|
||||
r.StopReason = "tool_use"
|
||||
}
|
||||
return r, nil
|
||||
}
|
||||
|
||||
// Next performs one chat-completions call.
|
||||
func (o *OpenAICompat) Next(ctx context.Context, req Request) (Reply, error) {
|
||||
body := toWire(req)
|
||||
body.Model = o.Model
|
||||
payload, err := json.Marshal(body)
|
||||
if err != nil {
|
||||
return Reply{}, fmt.Errorf("encode request: %w", err)
|
||||
}
|
||||
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, o.BaseURL+"/chat/completions", bytes.NewReader(payload))
|
||||
if err != nil {
|
||||
return Reply{}, fmt.Errorf("build request: %w", err)
|
||||
}
|
||||
httpReq.Header.Set("Authorization", "Bearer "+o.APIKey)
|
||||
httpReq.Header.Set("Content-Type", "application/json")
|
||||
|
||||
res, err := o.HTTP.Do(httpReq)
|
||||
if err != nil {
|
||||
return Reply{}, fmt.Errorf("model request: %w", err)
|
||||
}
|
||||
defer res.Body.Close()
|
||||
raw, _ := io.ReadAll(io.LimitReader(res.Body, 4<<20))
|
||||
|
||||
var out oaResponse
|
||||
_ = json.Unmarshal(raw, &out)
|
||||
if res.StatusCode/100 != 2 {
|
||||
msg := strings.TrimSpace(string(raw))
|
||||
if out.Error != nil && out.Error.Message != "" {
|
||||
msg = out.Error.Message
|
||||
}
|
||||
if len(msg) > 300 {
|
||||
msg = msg[:300]
|
||||
}
|
||||
return Reply{}, &ProviderError{Status: res.StatusCode, Message: msg}
|
||||
}
|
||||
return fromWire(out)
|
||||
}
|
||||
|
||||
// ProviderError is a non-2xx answer from the model provider. The status lets
|
||||
// the controller tell "rate limited" (429, common on free tiers) from a
|
||||
// genuine failure. The message never contains the API key.
|
||||
type ProviderError struct {
|
||||
Status int
|
||||
Message string
|
||||
}
|
||||
|
||||
func (e *ProviderError) Error() string {
|
||||
return fmt.Sprintf("model provider answered %d: %s", e.Status, e.Message)
|
||||
}
|
||||
130
internal/ai/playground/openai_compat_test.go
Normal file
130
internal/ai/playground/openai_compat_test.go
Normal file
@@ -0,0 +1,130 @@
|
||||
package playground
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// A fake OpenAI-compatible server: first call asks for a tool, second answers.
|
||||
func fakeProvider(t *testing.T, seen *[]map[string]any) *httptest.Server {
|
||||
t.Helper()
|
||||
calls := 0
|
||||
return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != "/openai/v1/chat/completions" {
|
||||
t.Errorf("path = %s", r.URL.Path)
|
||||
}
|
||||
if r.Header.Get("Authorization") != "Bearer test-key" {
|
||||
t.Errorf("auth header = %q", r.Header.Get("Authorization"))
|
||||
}
|
||||
b, _ := io.ReadAll(r.Body)
|
||||
var body map[string]any
|
||||
_ = json.Unmarshal(b, &body)
|
||||
*seen = append(*seen, body)
|
||||
calls++
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
if calls == 1 {
|
||||
_, _ = io.WriteString(w, `{"choices":[{"message":{"role":"assistant","content":null,"tool_calls":[
|
||||
{"id":"call_1","type":"function","function":{"name":"get_booking_cache","arguments":"{\"booking_id\":5}"}},
|
||||
{"id":"call_2","type":"function","function":{"name":"reassign_booking","arguments":"not json"}}]},
|
||||
"finish_reason":"tool_calls"}],"usage":{"prompt_tokens":100,"completion_tokens":20}}`)
|
||||
return
|
||||
}
|
||||
_, _ = io.WriteString(w, `{"choices":[{"message":{"role":"assistant","content":"Proposed a reassign."},"finish_reason":"stop"}],
|
||||
"usage":{"prompt_tokens":150,"completion_tokens":10}}`)
|
||||
}))
|
||||
}
|
||||
|
||||
func TestOpenAICompatRunsTheToolLoop(t *testing.T) {
|
||||
var seen []map[string]any
|
||||
srv := fakeProvider(t, &seen)
|
||||
defer srv.Close()
|
||||
|
||||
m := NewOpenAICompat(srv.URL+"/openai/v1/", "test-key", "openai/gpt-oss-120b")
|
||||
plan := mustPlan(t, "EXCEPTION_AGENT", "stall_response")
|
||||
plan.Model = m.ModelName()
|
||||
execs := map[string]Executor{
|
||||
"get_booking_cache": func(context.Context, json.RawMessage) (any, error) {
|
||||
return map[string]any{"bookingid": 5, "customername": "Ravi"}, nil
|
||||
},
|
||||
}
|
||||
tr, err := Run(context.Background(), m, plan, "booking 5 is stuck", execs)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if tr.Final != "Proposed a reassign." || tr.Turns != 2 || tr.Inputtokens != 250 || tr.Outputtokens != 30 {
|
||||
t.Fatalf("trace = %+v", tr)
|
||||
}
|
||||
if tr.Model != "openai/gpt-oss-120b" {
|
||||
t.Fatalf("model = %s", tr.Model)
|
||||
}
|
||||
if tr.Steps[0].Outcome != OutcomeExecuted || tr.Steps[1].Outcome != OutcomeProposed || string(tr.Steps[1].Input) != "{}" {
|
||||
t.Fatalf("steps = %+v", tr.Steps)
|
||||
}
|
||||
|
||||
// First request: system + user, the model from config, tools as functions
|
||||
// with object schemas.
|
||||
first := seen[0]
|
||||
if first["model"] != "openai/gpt-oss-120b" || first["max_completion_tokens"].(float64) != MaxTokens {
|
||||
t.Fatalf("first request = %v", first)
|
||||
}
|
||||
msgs := first["messages"].([]any)
|
||||
if msgs[0].(map[string]any)["role"] != "system" || msgs[1].(map[string]any)["role"] != "user" {
|
||||
t.Fatalf("messages = %v", msgs)
|
||||
}
|
||||
tool := first["tools"].([]any)[0].(map[string]any)
|
||||
params := tool["function"].(map[string]any)["parameters"].(map[string]any)
|
||||
if tool["type"] != "function" || params["type"] != "object" || params["properties"] == nil {
|
||||
t.Fatalf("tool = %v", tool)
|
||||
}
|
||||
|
||||
// Second request: the assistant's tool_calls echoed, then one tool message
|
||||
// per call, redacted.
|
||||
msgs = seen[1]["messages"].([]any)
|
||||
asst := msgs[2].(map[string]any)
|
||||
if asst["role"] != "assistant" || len(asst["tool_calls"].([]any)) != 2 {
|
||||
t.Fatalf("assistant echo = %v", asst)
|
||||
}
|
||||
res1, res2 := msgs[3].(map[string]any), msgs[4].(map[string]any)
|
||||
if res1["role"] != "tool" || res1["tool_call_id"] != "call_1" || res2["tool_call_id"] != "call_2" {
|
||||
t.Fatalf("tool results = %v %v", res1, res2)
|
||||
}
|
||||
if strings.Contains(res1["content"].(string), "Ravi") {
|
||||
t.Fatal("personal data reached the provider")
|
||||
}
|
||||
if !strings.Contains(res2["content"].(string), `"executed":false`) {
|
||||
t.Fatalf("write tool was not answered as a proposal: %v", res2["content"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAICompatReportsProviderErrors(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.WriteHeader(http.StatusTooManyRequests)
|
||||
_, _ = io.WriteString(w, `{"error":{"message":"Rate limit reached for model"}}`)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
_, err := NewOpenAICompat(srv.URL, "secret-key-value", "m").Next(context.Background(), Request{Turns: []Turn{{Role: "user", Text: "hi"}}})
|
||||
var pe *ProviderError
|
||||
if !errors.As(err, &pe) || pe.Status != 429 || !strings.Contains(pe.Message, "Rate limit") {
|
||||
t.Fatalf("err = %v", err)
|
||||
}
|
||||
if strings.Contains(err.Error(), "secret-key-value") {
|
||||
t.Fatal("the API key leaked into the error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAICompatStopWithToolCallsStillLoops(t *testing.T) {
|
||||
r, err := fromWire(oaResponse{Choices: []struct {
|
||||
Message oaMessage `json:"message"`
|
||||
FinishReason string `json:"finish_reason"`
|
||||
}{{Message: oaMessage{ToolCalls: []oaToolCall{{ID: "c", Function: oaFunctionCall{Name: "x", Arguments: "{}"}}}}, FinishReason: "stop"}}})
|
||||
if err != nil || r.StopReason != "tool_use" {
|
||||
t.Fatalf("reply = %+v, err = %v", r, err)
|
||||
}
|
||||
}
|
||||
341
internal/ai/playground/playground.go
Normal file
341
internal/ai/playground/playground.go
Normal file
@@ -0,0 +1,341 @@
|
||||
// Package playground runs one operator prompt through Claude with a registry
|
||||
// skill's tools, for Agent Studio's Test tab (Phase 6 of
|
||||
// krow_talent_app/docs/agent-platform-plan.md).
|
||||
//
|
||||
// The rules that make it safe to point at production:
|
||||
// - read tools the backend can serve run for real, and their results are
|
||||
// redacted (names, phones, addresses, emails, free text) before Claude sees
|
||||
// them — see redact.go;
|
||||
// - write, notify and event tools NEVER run: the call is answered with a
|
||||
// proposal and shown in the trace as "proposed";
|
||||
// - a tool outside the selected skill is refused;
|
||||
// - the loop is bounded (MaxTurns) and so is every tool call (ToolTimeout).
|
||||
//
|
||||
// Claude is reached through the Model interface, so the loop is tested with a
|
||||
// fake and the server wires in a real client only when one is configured.
|
||||
package playground
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"doormile/internal/ai/registry"
|
||||
)
|
||||
|
||||
// DefaultModel is used when the agent has no model pinned in the registry.
|
||||
const DefaultModel = "claude-opus-5-5"
|
||||
|
||||
const (
|
||||
MaxTurns = 6
|
||||
MaxTokens = 4096
|
||||
MaxPromptChars = 2000
|
||||
ToolTimeout = 5 * time.Second
|
||||
maxResultBytes = 16 * 1024
|
||||
)
|
||||
|
||||
// Tool kinds, as the registry stores them.
|
||||
const (
|
||||
kindRead = "read"
|
||||
)
|
||||
|
||||
// Outcomes of a tool call, as the trace shows them.
|
||||
const (
|
||||
OutcomeExecuted = "executed"
|
||||
OutcomeProposed = "proposed"
|
||||
OutcomeUnavailable = "unavailable"
|
||||
OutcomeError = "error"
|
||||
OutcomeRejected = "rejected"
|
||||
)
|
||||
|
||||
// ── The model boundary ──────────────────────────────────────────────────────
|
||||
|
||||
// ToolUse is a tool call Claude asked for.
|
||||
type ToolUse struct {
|
||||
ID string
|
||||
Name string
|
||||
Input json.RawMessage
|
||||
}
|
||||
|
||||
// Block is one content block of a reply. Text and tool_use are read by the
|
||||
// loop; anything else (thinking) is carried in Raw and sent back unchanged,
|
||||
// which the API requires within a tool-use turn.
|
||||
type Block struct {
|
||||
Type string
|
||||
Text string
|
||||
ToolUse *ToolUse
|
||||
Raw json.RawMessage
|
||||
}
|
||||
|
||||
// Reply is one model response.
|
||||
type Reply struct {
|
||||
Blocks []Block
|
||||
StopReason string
|
||||
InputTokens int64
|
||||
OutputTokens int64
|
||||
}
|
||||
|
||||
// ToolResult answers one ToolUse.
|
||||
type ToolResult struct {
|
||||
ToolUseID string
|
||||
Content string
|
||||
IsError bool
|
||||
}
|
||||
|
||||
// Turn is one message of the conversation: the user's prompt, an assistant
|
||||
// reply, or the user turn carrying tool results.
|
||||
type Turn struct {
|
||||
Role string // "user" or "assistant"
|
||||
Text string
|
||||
Assistant []Block
|
||||
Results []ToolResult
|
||||
}
|
||||
|
||||
// ToolDef is a tool as offered to the model.
|
||||
type ToolDef struct {
|
||||
Name string
|
||||
Description string
|
||||
InputSchema json.RawMessage
|
||||
}
|
||||
|
||||
// Request is one model call.
|
||||
type Request struct {
|
||||
Model string
|
||||
System string
|
||||
MaxTokens int64
|
||||
Tools []ToolDef
|
||||
Turns []Turn
|
||||
}
|
||||
|
||||
// Model is the one call the loop needs from Claude.
|
||||
type Model interface {
|
||||
Next(ctx context.Context, req Request) (Reply, error)
|
||||
}
|
||||
|
||||
// ── Plan: what a run may use ────────────────────────────────────────────────
|
||||
|
||||
// ErrNotFound is returned when the agent or skill does not exist.
|
||||
var ErrNotFound = errors.New("not found")
|
||||
|
||||
// Plan is the resolved agent, skill and tools for one run.
|
||||
type Plan struct {
|
||||
AgentID string
|
||||
AgentName string
|
||||
SkillID string
|
||||
Model string
|
||||
System string
|
||||
Tools []ToolDef
|
||||
kinds map[string]string
|
||||
}
|
||||
|
||||
// Prepare resolves a run from the registry. skillID may be empty: the run then
|
||||
// gets every tool of the agent's enabled skills.
|
||||
func Prepare(snap *registry.Snapshot, agentID, skillID string) (*Plan, error) {
|
||||
var agent *registry.AgentView
|
||||
for i := range snap.Agents {
|
||||
if snap.Agents[i].Agentid == agentID {
|
||||
agent = &snap.Agents[i]
|
||||
}
|
||||
}
|
||||
if agent == nil {
|
||||
return nil, fmt.Errorf("agent %q: %w", agentID, ErrNotFound)
|
||||
}
|
||||
|
||||
toolNames := map[string]bool{}
|
||||
var skillLines []string
|
||||
found := skillID == ""
|
||||
for _, s := range snap.Skills {
|
||||
if s.Agentid != agentID {
|
||||
continue
|
||||
}
|
||||
if skillID != "" && s.Skillid != skillID {
|
||||
continue
|
||||
}
|
||||
if skillID == "" && !s.Enabled {
|
||||
continue
|
||||
}
|
||||
found = true
|
||||
skillLines = append(skillLines, fmt.Sprintf("- %s: %s", s.Title, s.Description))
|
||||
for _, t := range s.Tools {
|
||||
toolNames[t] = true
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
return nil, fmt.Errorf("skill %q of agent %q: %w", skillID, agentID, ErrNotFound)
|
||||
}
|
||||
|
||||
plan := &Plan{
|
||||
AgentID: agent.Agentid, AgentName: agent.Name, SkillID: skillID,
|
||||
Model: agent.Model, kinds: map[string]string{}, Tools: []ToolDef{},
|
||||
}
|
||||
if plan.Model == "" {
|
||||
plan.Model = DefaultModel
|
||||
}
|
||||
for _, t := range snap.Tools {
|
||||
if !toolNames[t.Toolname] {
|
||||
continue
|
||||
}
|
||||
desc := t.Description
|
||||
if t.Kind != kindRead {
|
||||
desc += " [Playground: NOT executed — calling it records a proposal for a human.]"
|
||||
}
|
||||
plan.Tools = append(plan.Tools, ToolDef{Name: t.Toolname, Description: desc, InputSchema: t.Inputschema})
|
||||
plan.kinds[t.Toolname] = t.Kind
|
||||
}
|
||||
|
||||
plan.System = strings.Join([]string{
|
||||
fmt.Sprintf("You are %s, an agent in Doormile's delivery operations, being tested by an operator in the Agent Studio playground.", agent.Name),
|
||||
"Purpose: " + agent.Purpose,
|
||||
"Skills in scope:\n" + strings.Join(skillLines, "\n"),
|
||||
"Read tools return live data with personal details (names, phones, addresses, notes) removed; do not ask for them.",
|
||||
"Write, notify and event tools are not executed here: calling one records a proposal for a human to review. Say plainly what you would do and why.",
|
||||
"If a tool is unavailable, say so rather than guessing its result. Keep the final answer short and concrete.",
|
||||
}, "\n\n")
|
||||
return plan, nil
|
||||
}
|
||||
|
||||
// ── The run ─────────────────────────────────────────────────────────────────
|
||||
|
||||
// Executor runs one read tool. Its result is redacted before the model sees it.
|
||||
type Executor func(ctx context.Context, input json.RawMessage) (any, error)
|
||||
|
||||
// Step is one line of the trace the console shows.
|
||||
type Step struct {
|
||||
Kind string `json:"kind"` // "text" or "tool"
|
||||
Text string `json:"text,omitempty"`
|
||||
Tool string `json:"tool,omitempty"`
|
||||
Toolkind string `json:"toolkind,omitempty"`
|
||||
Input json.RawMessage `json:"input,omitempty"`
|
||||
Outcome string `json:"outcome,omitempty"`
|
||||
Result json.RawMessage `json:"result,omitempty"`
|
||||
Ms int64 `json:"ms"`
|
||||
}
|
||||
|
||||
// Trace is the whole run.
|
||||
type Trace struct {
|
||||
Agentid string `json:"agentid"`
|
||||
Skillid string `json:"skillid"`
|
||||
Model string `json:"model"`
|
||||
Steps []Step `json:"steps"`
|
||||
Final string `json:"final"`
|
||||
Stopreason string `json:"stopreason"`
|
||||
Turns int `json:"turns"`
|
||||
Inputtokens int64 `json:"inputtokens"`
|
||||
Outputtokens int64 `json:"outputtokens"`
|
||||
Ms int64 `json:"ms"`
|
||||
}
|
||||
|
||||
// Run executes the prompt. A model error ends the run with the error; the
|
||||
// trace so far is still returned so the console can show how far it got.
|
||||
func Run(ctx context.Context, m Model, plan *Plan, prompt string, execs map[string]Executor) (*Trace, error) {
|
||||
started := time.Now()
|
||||
tr := &Trace{Agentid: plan.AgentID, Skillid: plan.SkillID, Model: plan.Model, Steps: []Step{}}
|
||||
turns := []Turn{{Role: "user", Text: prompt}}
|
||||
|
||||
defer func() { tr.Ms = time.Since(started).Milliseconds() }()
|
||||
|
||||
for tr.Turns < MaxTurns {
|
||||
callStarted := time.Now()
|
||||
reply, err := m.Next(ctx, Request{
|
||||
Model: plan.Model, System: plan.System, MaxTokens: MaxTokens, Tools: plan.Tools, Turns: turns,
|
||||
})
|
||||
tr.Turns++
|
||||
if err != nil {
|
||||
tr.Stopreason = "error"
|
||||
return tr, err
|
||||
}
|
||||
tr.Inputtokens += reply.InputTokens
|
||||
tr.Outputtokens += reply.OutputTokens
|
||||
tr.Stopreason = reply.StopReason
|
||||
modelMs := time.Since(callStarted).Milliseconds()
|
||||
|
||||
turns = append(turns, Turn{Role: "assistant", Assistant: reply.Blocks})
|
||||
|
||||
var texts []string
|
||||
var results []ToolResult
|
||||
for _, b := range reply.Blocks {
|
||||
switch {
|
||||
case b.Type == "text" && strings.TrimSpace(b.Text) != "":
|
||||
texts = append(texts, b.Text)
|
||||
tr.Steps = append(tr.Steps, Step{Kind: "text", Text: b.Text, Ms: modelMs})
|
||||
modelMs = 0
|
||||
case b.Type == "tool_use" && b.ToolUse != nil:
|
||||
step, res := callTool(ctx, plan, execs, *b.ToolUse)
|
||||
tr.Steps = append(tr.Steps, step)
|
||||
results = append(results, res)
|
||||
}
|
||||
}
|
||||
|
||||
if reply.StopReason != "tool_use" || len(results) == 0 {
|
||||
tr.Final = strings.Join(texts, "\n\n")
|
||||
return tr, nil
|
||||
}
|
||||
turns = append(turns, Turn{Role: "user", Results: results})
|
||||
}
|
||||
|
||||
tr.Stopreason = "max_turns"
|
||||
return tr, nil
|
||||
}
|
||||
|
||||
func callTool(ctx context.Context, plan *Plan, execs map[string]Executor, use ToolUse) (Step, ToolResult) {
|
||||
started := time.Now()
|
||||
input := use.Input
|
||||
if len(input) == 0 || !json.Valid(input) {
|
||||
input = json.RawMessage(`{}`)
|
||||
}
|
||||
kind, inSkill := plan.kinds[use.Name]
|
||||
step := Step{Kind: "tool", Tool: use.Name, Toolkind: kind, Input: input}
|
||||
res := ToolResult{ToolUseID: use.ID}
|
||||
|
||||
finish := func(outcome string, payload any, isError bool) (Step, ToolResult) {
|
||||
body := encode(payload)
|
||||
step.Outcome, step.Result, step.Ms = outcome, body, time.Since(started).Milliseconds()
|
||||
res.Content, res.IsError = string(body), isError
|
||||
return step, res
|
||||
}
|
||||
|
||||
switch {
|
||||
case !inSkill:
|
||||
return finish(OutcomeRejected, map[string]string{"error": "This tool is not part of the selected skill."}, true)
|
||||
|
||||
case kind != kindRead:
|
||||
return finish(OutcomeProposed, map[string]any{
|
||||
"executed": false,
|
||||
"proposal": map[string]any{"tool": use.Name, "input": input},
|
||||
"note": "Playground: recorded as a proposal for a human. Nothing was changed.",
|
||||
}, false)
|
||||
|
||||
case execs[use.Name] == nil:
|
||||
return finish(OutcomeUnavailable, map[string]string{
|
||||
"error": "This read tool is not available in the playground (it runs inside AI_engine or calls an external service).",
|
||||
}, true)
|
||||
}
|
||||
|
||||
tctx, cancel := context.WithTimeout(ctx, ToolTimeout)
|
||||
defer cancel()
|
||||
out, err := execs[use.Name](tctx, input)
|
||||
if err != nil {
|
||||
return finish(OutcomeError, map[string]string{"error": err.Error()}, true)
|
||||
}
|
||||
return finish(OutcomeExecuted, Redact(out), false)
|
||||
}
|
||||
|
||||
// encode marshals a tool result, capped so one large read cannot blow the
|
||||
// context window or the console.
|
||||
func encode(v any) json.RawMessage {
|
||||
b, err := json.Marshal(v)
|
||||
if err != nil {
|
||||
b, _ = json.Marshal(map[string]string{"error": "result could not be encoded"})
|
||||
}
|
||||
if len(b) > maxResultBytes {
|
||||
b, _ = json.Marshal(map[string]any{
|
||||
"truncated": true,
|
||||
"note": fmt.Sprintf("Result was %d bytes; showing the first %d.", len(b), maxResultBytes),
|
||||
"partial": string(b[:maxResultBytes]),
|
||||
})
|
||||
}
|
||||
return b
|
||||
}
|
||||
280
internal/ai/playground/playground_test.go
Normal file
280
internal/ai/playground/playground_test.go
Normal file
@@ -0,0 +1,280 @@
|
||||
package playground
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"doormile/internal/ai/registry"
|
||||
"doormile/models"
|
||||
)
|
||||
|
||||
// fakeModel replays scripted replies and records every request.
|
||||
type fakeModel struct {
|
||||
replies []Reply
|
||||
err error
|
||||
reqs []Request
|
||||
}
|
||||
|
||||
func (f *fakeModel) Next(_ context.Context, req Request) (Reply, error) {
|
||||
f.reqs = append(f.reqs, req)
|
||||
if f.err != nil {
|
||||
return Reply{}, f.err
|
||||
}
|
||||
if len(f.replies) == 0 {
|
||||
return Reply{StopReason: "end_turn", Blocks: []Block{{Type: "text", Text: "done"}}}, nil
|
||||
}
|
||||
r := f.replies[0]
|
||||
f.replies = f.replies[1:]
|
||||
return r, nil
|
||||
}
|
||||
|
||||
func toolCall(id, name, input string) Block {
|
||||
return Block{Type: "tool_use", ToolUse: &ToolUse{ID: id, Name: name, Input: json.RawMessage(input)}}
|
||||
}
|
||||
|
||||
func testSnapshot() *registry.Snapshot {
|
||||
agents := []models.AIAgent{
|
||||
{Agentid: "EXCEPTION_AGENT", Name: "Exception", Purpose: "Handles stalled riders.", Model: "claude-sonnet-5-5"},
|
||||
{Agentid: "CONSOLE_OPS_AGENT", Name: "Console Ops Agent", Purpose: "Watches the board."},
|
||||
}
|
||||
tools := []models.AITool{
|
||||
{Toolname: "get_booking_cache", Kind: "read", Description: "Read a booking.", Inputschema: `{"type":"object"}`},
|
||||
{Toolname: "nearby_milers", Kind: "read", Description: "Riders near a point.", Inputschema: `{"type":"object"}`},
|
||||
{Toolname: "decide_stall_response", Kind: "read", Description: "Engine decision.", Inputschema: `{"type":"object"}`},
|
||||
{Toolname: "reassign_booking", Kind: "write", Description: "Reassign.", Inputschema: `{"type":"object"}`},
|
||||
{Toolname: "scan_bookings", Kind: "read", Description: "Scan.", Inputschema: `{"type":"object"}`},
|
||||
}
|
||||
skills := []models.AISkill{
|
||||
{Skillid: "stall_response", Agentid: "EXCEPTION_AGENT", Title: "Stall", Description: "Respond to stalls.", Enabled: true},
|
||||
{Skillid: "off_skill", Agentid: "EXCEPTION_AGENT", Title: "Off", Enabled: false},
|
||||
{Skillid: "late_dispatch", Agentid: "CONSOLE_OPS_AGENT", Title: "Late", Enabled: true},
|
||||
}
|
||||
links := []models.AISkillTool{
|
||||
{Skillid: "stall_response", Toolname: "get_booking_cache"},
|
||||
{Skillid: "stall_response", Toolname: "nearby_milers"},
|
||||
{Skillid: "stall_response", Toolname: "decide_stall_response"},
|
||||
{Skillid: "stall_response", Toolname: "reassign_booking"},
|
||||
{Skillid: "off_skill", Toolname: "scan_bookings"},
|
||||
{Skillid: "late_dispatch", Toolname: "scan_bookings"},
|
||||
}
|
||||
return registry.Build(agents, tools, skills, links)
|
||||
}
|
||||
|
||||
func mustPlan(t *testing.T, agent, skill string) *Plan {
|
||||
t.Helper()
|
||||
p, err := Prepare(testSnapshot(), agent, skill)
|
||||
if err != nil {
|
||||
t.Fatalf("Prepare: %v", err)
|
||||
}
|
||||
return p
|
||||
}
|
||||
|
||||
func toolNames(p *Plan) []string {
|
||||
var out []string
|
||||
for _, t := range p.Tools {
|
||||
out = append(out, t.Name)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// ── Prepare ─────────────────────────────────────────────────────────────────
|
||||
|
||||
func TestPrepareUsesSkillToolsAndPinnedModel(t *testing.T) {
|
||||
p := mustPlan(t, "EXCEPTION_AGENT", "stall_response")
|
||||
if p.Model != "claude-sonnet-5-5" {
|
||||
t.Fatalf("model = %q, want the registry pin", p.Model)
|
||||
}
|
||||
got := strings.Join(toolNames(p), ",")
|
||||
// Registry order (Load sorts by toolname; this fixture is in its own order).
|
||||
if got != "get_booking_cache,nearby_milers,decide_stall_response,reassign_booking" {
|
||||
t.Fatalf("tools = %s", got)
|
||||
}
|
||||
for _, td := range p.Tools {
|
||||
marked := strings.Contains(td.Description, "NOT executed")
|
||||
if (td.Name == "reassign_booking") != marked {
|
||||
t.Fatalf("%s: write-tool marking = %v", td.Name, marked)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrepareDefaultsModelAndSkipsDisabledSkillsWhenNoSkillGiven(t *testing.T) {
|
||||
p := mustPlan(t, "CONSOLE_OPS_AGENT", "")
|
||||
if p.Model != DefaultModel {
|
||||
t.Fatalf("model = %q, want %q", p.Model, DefaultModel)
|
||||
}
|
||||
p = mustPlan(t, "EXCEPTION_AGENT", "")
|
||||
for _, n := range toolNames(p) {
|
||||
if n == "scan_bookings" {
|
||||
t.Fatal("a disabled skill's tool was offered")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrepareNotFound(t *testing.T) {
|
||||
for _, c := range [][2]string{{"NOPE", ""}, {"EXCEPTION_AGENT", "late_dispatch"}, {"EXCEPTION_AGENT", "missing"}} {
|
||||
if _, err := Prepare(testSnapshot(), c[0], c[1]); !errors.Is(err, ErrNotFound) {
|
||||
t.Fatalf("%v: err = %v, want ErrNotFound", c, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ── Run ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
func TestRunTextOnly(t *testing.T) {
|
||||
m := &fakeModel{replies: []Reply{{StopReason: "end_turn", InputTokens: 10, OutputTokens: 5,
|
||||
Blocks: []Block{{Type: "thinking", Raw: json.RawMessage(`{"type":"thinking"}`)}, {Type: "text", Text: "All clear."}}}}}
|
||||
tr, err := Run(context.Background(), m, mustPlan(t, "EXCEPTION_AGENT", "stall_response"), "status?", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if tr.Final != "All clear." || tr.Turns != 1 || tr.Inputtokens != 10 || tr.Outputtokens != 5 || tr.Model != "claude-sonnet-5-5" {
|
||||
t.Fatalf("trace = %+v", tr)
|
||||
}
|
||||
if len(tr.Steps) != 1 || tr.Steps[0].Kind != "text" {
|
||||
t.Fatalf("steps = %+v", tr.Steps)
|
||||
}
|
||||
req := m.reqs[0]
|
||||
if req.Turns[0].Text != "status?" || req.MaxTokens != MaxTokens || len(req.Tools) != 4 || req.System == "" {
|
||||
t.Fatalf("request = %+v", req)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunToolOutcomes(t *testing.T) {
|
||||
executed := 0
|
||||
execs := map[string]Executor{
|
||||
"get_booking_cache": func(_ context.Context, in json.RawMessage) (any, error) {
|
||||
executed++
|
||||
return map[string]any{"bookingid": 5, "customername": "Ravi", "notes": "call 9876543210"}, nil
|
||||
},
|
||||
// Present but NOT in the skill — must never run.
|
||||
"scan_bookings": func(context.Context, json.RawMessage) (any, error) {
|
||||
t.Fatal("a tool outside the skill was executed")
|
||||
return nil, nil
|
||||
},
|
||||
"nearby_milers": func(context.Context, json.RawMessage) (any, error) { return nil, errors.New("positions unavailable") },
|
||||
}
|
||||
m := &fakeModel{replies: []Reply{
|
||||
{StopReason: "tool_use", Blocks: []Block{
|
||||
toolCall("t1", "get_booking_cache", `{"booking_id":5}`),
|
||||
toolCall("t2", "reassign_booking", `{"booking_id":5}`),
|
||||
toolCall("t3", "scan_bookings", `{}`),
|
||||
toolCall("t4", "decide_stall_response", `{}`),
|
||||
toolCall("t5", "nearby_milers", `{"lat":11,"lon":77}`),
|
||||
}},
|
||||
{StopReason: "end_turn", Blocks: []Block{{Type: "text", Text: "Proposed a reassign."}}},
|
||||
}}
|
||||
tr, err := Run(context.Background(), m, mustPlan(t, "EXCEPTION_AGENT", "stall_response"), "booking 5 is stuck", execs)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
want := map[string]string{
|
||||
"get_booking_cache": OutcomeExecuted, "reassign_booking": OutcomeProposed, "scan_bookings": OutcomeRejected,
|
||||
"decide_stall_response": OutcomeUnavailable, "nearby_milers": OutcomeError,
|
||||
}
|
||||
for _, s := range tr.Steps {
|
||||
if s.Kind == "tool" && want[s.Tool] != s.Outcome {
|
||||
t.Fatalf("%s outcome = %s, want %s", s.Tool, s.Outcome, want[s.Tool])
|
||||
}
|
||||
}
|
||||
if executed != 1 {
|
||||
t.Fatalf("read tool executed %d times", executed)
|
||||
}
|
||||
|
||||
// The second request carries all five results, in order, and redacted.
|
||||
results := m.reqs[1].Turns[2].Results
|
||||
if len(results) != 5 || results[0].ToolUseID != "t1" {
|
||||
t.Fatalf("results = %+v", results)
|
||||
}
|
||||
if strings.Contains(results[0].Content, "Ravi") || strings.Contains(results[0].Content, "9876543210") {
|
||||
t.Fatalf("personal data reached the model: %s", results[0].Content)
|
||||
}
|
||||
if !strings.Contains(results[1].Content, `"executed":false`) || results[1].IsError {
|
||||
t.Fatalf("write tool result = %+v", results[1])
|
||||
}
|
||||
for _, i := range []int{2, 3, 4} {
|
||||
if !results[i].IsError {
|
||||
t.Fatalf("result %d should be an error: %+v", i, results[i])
|
||||
}
|
||||
}
|
||||
// The assistant turn (with its tool_use blocks) is echoed back before the results.
|
||||
if m.reqs[1].Turns[1].Role != "assistant" || len(m.reqs[1].Turns[1].Assistant) != 5 {
|
||||
t.Fatalf("assistant turn not echoed: %+v", m.reqs[1].Turns[1])
|
||||
}
|
||||
if tr.Final != "Proposed a reassign." || tr.Turns != 2 {
|
||||
t.Fatalf("trace = %+v", tr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunStopsAtMaxTurns(t *testing.T) {
|
||||
var replies []Reply
|
||||
for i := 0; i < MaxTurns+2; i++ {
|
||||
replies = append(replies, Reply{StopReason: "tool_use", Blocks: []Block{toolCall("t", "reassign_booking", `{}`)}})
|
||||
}
|
||||
m := &fakeModel{replies: replies}
|
||||
tr, err := Run(context.Background(), m, mustPlan(t, "EXCEPTION_AGENT", "stall_response"), "loop", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if tr.Turns != MaxTurns || tr.Stopreason != "max_turns" || len(m.reqs) != MaxTurns {
|
||||
t.Fatalf("turns = %d, stop = %s, calls = %d", tr.Turns, tr.Stopreason, len(m.reqs))
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunModelError(t *testing.T) {
|
||||
tr, err := Run(context.Background(), &fakeModel{err: errors.New("overloaded")}, mustPlan(t, "EXCEPTION_AGENT", "stall_response"), "x", nil)
|
||||
if err == nil || tr == nil || tr.Stopreason != "error" {
|
||||
t.Fatalf("err = %v, trace = %+v", err, tr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInvalidToolInputBecomesEmptyObject(t *testing.T) {
|
||||
m := &fakeModel{replies: []Reply{{StopReason: "tool_use", Blocks: []Block{toolCall("t", "reassign_booking", `{not json`)}}}}
|
||||
tr, _ := Run(context.Background(), m, mustPlan(t, "EXCEPTION_AGENT", "stall_response"), "x", nil)
|
||||
if string(tr.Steps[0].Input) != "{}" {
|
||||
t.Fatalf("input = %s", tr.Steps[0].Input)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEncodeCapsLargeResults(t *testing.T) {
|
||||
b := encode(map[string]string{"blob": strings.Repeat("x", maxResultBytes*2)})
|
||||
if len(b) > maxResultBytes+1024 || !strings.Contains(string(b), `"truncated":true`) {
|
||||
t.Fatalf("len = %d", len(b))
|
||||
}
|
||||
}
|
||||
|
||||
// ── Redact ──────────────────────────────────────────────────────────────────
|
||||
|
||||
func TestRedact(t *testing.T) {
|
||||
in := map[string]any{
|
||||
"bookingid": 7,
|
||||
"customerName": "Ravi",
|
||||
"pickupaddress": "12 MG Road",
|
||||
"status": "Created",
|
||||
"createdat_ist": "2026-09-29 12:30",
|
||||
"deliverycity": "Coimbatore",
|
||||
"cancelreason": "customer asked",
|
||||
"nested": []any{map[string]any{"phone": "9876543210", "hub": "call +91 98765 43210 or a@b.co"}},
|
||||
"missingnote": nil,
|
||||
}
|
||||
out := Redact(in).(map[string]any)
|
||||
for _, k := range []string{"customerName", "pickupaddress", "cancelreason"} {
|
||||
if out[k] != Redacted {
|
||||
t.Fatalf("%s = %v, want redacted", k, out[k])
|
||||
}
|
||||
}
|
||||
for k, want := range map[string]any{"status": "Created", "createdat_ist": "2026-09-29 12:30", "deliverycity": "Coimbatore", "bookingid": float64(7)} {
|
||||
if out[k] != want {
|
||||
t.Fatalf("%s = %v, want %v (must not be redacted)", k, out[k], want)
|
||||
}
|
||||
}
|
||||
nested := out["nested"].([]any)[0].(map[string]any)
|
||||
if nested["phone"] != Redacted || strings.ContainsAny(nested["hub"].(string), "@") || strings.Contains(nested["hub"].(string), "98765") {
|
||||
t.Fatalf("nested = %v", nested)
|
||||
}
|
||||
if out["missingnote"] != nil {
|
||||
t.Fatal("a null personal field should stay null")
|
||||
}
|
||||
}
|
||||
78
internal/ai/playground/redact.go
Normal file
78
internal/ai/playground/redact.go
Normal file
@@ -0,0 +1,78 @@
|
||||
package playground
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"regexp"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// Redacted replaces every value Redact removes.
|
||||
const Redacted = "[redacted]"
|
||||
|
||||
// piiKeys: a field whose name contains one of these is personal or free text
|
||||
// (free text is where people type phone numbers and addresses).
|
||||
var piiKeys = []string{
|
||||
"name", "phone", "mobile", "email", "address", "landmark", "otp", "contact",
|
||||
"note", "remark", "instruction", "description", "reason", "comment",
|
||||
}
|
||||
|
||||
var (
|
||||
emailLike = regexp.MustCompile(`[A-Za-z0-9._%+-]+@[A-Za-z0-9.-]+\.[A-Za-z]{2,}`)
|
||||
// An Indian mobile (10 digits from 6-9, optional +91). Narrow on purpose:
|
||||
// a looser digit-run pattern also masks dates like "2026-09-29 12".
|
||||
phoneLike = regexp.MustCompile(`(?:\+?91[\s-]?)?\b[6-9]\d{4}[\s-]?\d{5}\b`)
|
||||
)
|
||||
|
||||
// Redact returns a copy of v, as plain JSON values, with personal data
|
||||
// removed: values under personal keys are replaced, and any remaining string
|
||||
// that contains an email or a phone-like number has it masked.
|
||||
//
|
||||
// Executors already select only non-personal columns; this is the backstop
|
||||
// that makes a later column addition safe by default.
|
||||
func Redact(v any) any {
|
||||
b, err := json.Marshal(v)
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
var generic any
|
||||
if err := json.Unmarshal(b, &generic); err != nil {
|
||||
return nil
|
||||
}
|
||||
return redactValue(generic)
|
||||
}
|
||||
|
||||
func redactValue(v any) any {
|
||||
switch t := v.(type) {
|
||||
case map[string]any:
|
||||
out := make(map[string]any, len(t))
|
||||
for k, val := range t {
|
||||
if isPIIKey(k) && val != nil {
|
||||
out[k] = Redacted
|
||||
continue
|
||||
}
|
||||
out[k] = redactValue(val)
|
||||
}
|
||||
return out
|
||||
case []any:
|
||||
out := make([]any, len(t))
|
||||
for i, val := range t {
|
||||
out[i] = redactValue(val)
|
||||
}
|
||||
return out
|
||||
case string:
|
||||
s := emailLike.ReplaceAllString(t, Redacted)
|
||||
return phoneLike.ReplaceAllString(s, Redacted)
|
||||
default:
|
||||
return v
|
||||
}
|
||||
}
|
||||
|
||||
func isPIIKey(k string) bool {
|
||||
k = strings.ToLower(k)
|
||||
for _, p := range piiKeys {
|
||||
if strings.Contains(k, p) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
178
internal/ai/playground/tools.go
Normal file
178
internal/ai/playground/tools.go
Normal file
@@ -0,0 +1,178 @@
|
||||
package playground
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"math"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/redis/go-redis/v9"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// The read tools the backend serves itself. Every other read tool in the
|
||||
// registry runs inside AI_engine (decide_*), calls an external service
|
||||
// (sequence_stops, simulate_pricing_quote) or an engine-only internal route
|
||||
// (list_express_*), and answers "unavailable" in the playground.
|
||||
//
|
||||
// Each executor selects named, non-personal columns only — never addresses,
|
||||
// names, phones or notes — and Redact runs over the result as a backstop.
|
||||
// Coordinates are rounded to 2 decimals (about 1 km).
|
||||
|
||||
const (
|
||||
scanDefaultLimit = 20
|
||||
scanMaxLimit = 50
|
||||
nearbyMaxKm = 10.0
|
||||
nearbyMaxCount = 20
|
||||
)
|
||||
|
||||
// bookingRow is the only projection of a booking the playground exposes.
|
||||
type bookingRow struct {
|
||||
Bookingid int `json:"bookingid"`
|
||||
Bookingno string `json:"bookingno"`
|
||||
Tenantid *int `json:"tenantid"`
|
||||
Status string `json:"status"`
|
||||
Pickuppincode string `json:"pickuppincode"`
|
||||
Deliverypincode string `json:"deliverypincode"`
|
||||
Deliverycity string `json:"deliverycity"`
|
||||
Pickuplatitude float64 `json:"pickuplat"`
|
||||
Pickuplongitude float64 `json:"pickuplon"`
|
||||
Deliverylatitude float64 `json:"deliverylat"`
|
||||
Deliverylongitude float64 `json:"deliverylon"`
|
||||
Assignedmileruserid *int `json:"assignedmileruserid"`
|
||||
Routekm *float64 `json:"routekm"`
|
||||
Createdat *time.Time `json:"-"`
|
||||
Createdatist string `json:"createdat_ist,omitempty" gorm:"-"`
|
||||
}
|
||||
|
||||
const bookingColumns = "bookingid, bookingno, tenantid, status, pickuppincode, deliverypincode, deliverycity, " +
|
||||
"pickuplatitude, pickuplongitude, deliverylatitude, deliverylongitude, assignedmileruserid, routekm, createdat"
|
||||
|
||||
func round2(f float64) float64 { return math.Round(f*100) / 100 }
|
||||
|
||||
func (r *bookingRow) tidy() {
|
||||
r.Pickuplatitude, r.Pickuplongitude = round2(r.Pickuplatitude), round2(r.Pickuplongitude)
|
||||
r.Deliverylatitude, r.Deliverylongitude = round2(r.Deliverylatitude), round2(r.Deliverylongitude)
|
||||
if r.Createdat != nil {
|
||||
// pickupbookings.createdat is timestamp WITHOUT zone holding IST digits.
|
||||
r.Createdatist = r.Createdat.Format("2006-01-02 15:04")
|
||||
}
|
||||
}
|
||||
|
||||
// Executors returns the read tools this backend can serve. A nil db or rdb
|
||||
// just leaves the tools that need it out (they then answer "unavailable").
|
||||
func Executors(db *gorm.DB, rdb *redis.Client) map[string]Executor {
|
||||
execs := map[string]Executor{}
|
||||
if db != nil {
|
||||
execs["get_booking_cache"] = getBooking(db)
|
||||
execs["scan_bookings"] = scanBookings(db)
|
||||
}
|
||||
if rdb != nil {
|
||||
execs["nearby_milers"] = nearbyMilers(rdb)
|
||||
}
|
||||
return execs
|
||||
}
|
||||
|
||||
func decode(input json.RawMessage, into any) error {
|
||||
if len(input) == 0 {
|
||||
return nil
|
||||
}
|
||||
if err := json.Unmarshal(input, into); err != nil {
|
||||
return fmt.Errorf("invalid input: %v", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func getBooking(db *gorm.DB) Executor {
|
||||
return func(ctx context.Context, input json.RawMessage) (any, error) {
|
||||
var in struct {
|
||||
BookingID int `json:"booking_id"`
|
||||
}
|
||||
if err := decode(input, &in); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if in.BookingID <= 0 {
|
||||
return nil, errors.New("booking_id (a positive integer) is required")
|
||||
}
|
||||
var row bookingRow
|
||||
res := db.WithContext(ctx).Table("pickupbookings").Select(bookingColumns).
|
||||
Where("bookingid = ?", in.BookingID).Limit(1).Scan(&row)
|
||||
if res.Error != nil {
|
||||
return nil, errors.New("booking lookup failed")
|
||||
}
|
||||
if res.RowsAffected == 0 {
|
||||
return map[string]any{"found": false, "booking_id": in.BookingID}, nil
|
||||
}
|
||||
row.tidy()
|
||||
return map[string]any{"found": true, "booking": row, "source": "pickupbookings table"}, nil
|
||||
}
|
||||
}
|
||||
|
||||
func scanBookings(db *gorm.DB) Executor {
|
||||
return func(ctx context.Context, input json.RawMessage) (any, error) {
|
||||
var in struct {
|
||||
Status string `json:"status"`
|
||||
Limit int `json:"limit"`
|
||||
}
|
||||
if err := decode(input, &in); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if in.Limit <= 0 {
|
||||
in.Limit = scanDefaultLimit
|
||||
}
|
||||
if in.Limit > scanMaxLimit {
|
||||
in.Limit = scanMaxLimit
|
||||
}
|
||||
q := db.WithContext(ctx).Table("pickupbookings").Select(bookingColumns)
|
||||
if s := strings.TrimSpace(in.Status); s != "" {
|
||||
q = q.Where("status = ?", s)
|
||||
}
|
||||
var rows []bookingRow
|
||||
if err := q.Order("bookingid DESC").Limit(in.Limit).Scan(&rows).Error; err != nil {
|
||||
return nil, errors.New("booking scan failed")
|
||||
}
|
||||
byStatus := map[string]int{}
|
||||
for i := range rows {
|
||||
rows[i].tidy()
|
||||
byStatus[rows[i].Status]++
|
||||
}
|
||||
return map[string]any{"count": len(rows), "bystatus": byStatus, "bookings": rows, "order": "newest first"}, nil
|
||||
}
|
||||
}
|
||||
|
||||
func nearbyMilers(rdb *redis.Client) Executor {
|
||||
return func(ctx context.Context, input json.RawMessage) (any, error) {
|
||||
var in struct {
|
||||
Lat *float64 `json:"lat"`
|
||||
Lon *float64 `json:"lon"`
|
||||
RadiusKm float64 `json:"radius_km"`
|
||||
}
|
||||
if err := decode(input, &in); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if in.Lat == nil || in.Lon == nil || math.Abs(*in.Lat) > 90 || math.Abs(*in.Lon) > 180 {
|
||||
return nil, errors.New("lat and lon (decimal degrees) are required")
|
||||
}
|
||||
if in.RadiusKm <= 0 || in.RadiusKm > nearbyMaxKm {
|
||||
in.RadiusKm = 5
|
||||
}
|
||||
locs, err := rdb.GeoSearchLocation(ctx, "milers:locations", &redis.GeoSearchLocationQuery{
|
||||
GeoSearchQuery: redis.GeoSearchQuery{
|
||||
Longitude: *in.Lon, Latitude: *in.Lat,
|
||||
Radius: in.RadiusKm, RadiusUnit: "km", Sort: "ASC", Count: nearbyMaxCount,
|
||||
},
|
||||
WithDist: true,
|
||||
}).Result()
|
||||
if err != nil {
|
||||
return nil, errors.New("live rider positions are unavailable")
|
||||
}
|
||||
riders := make([]map[string]any, 0, len(locs))
|
||||
for _, l := range locs {
|
||||
riders = append(riders, map[string]any{"miler": l.Name, "distancekm": math.Round(l.Dist*100) / 100})
|
||||
}
|
||||
return map[string]any{"count": len(riders), "radiuskm": in.RadiusKm, "riders": riders}, nil
|
||||
}
|
||||
}
|
||||
82
internal/ai/playground/tools_integration_test.go
Normal file
82
internal/ai/playground/tools_integration_test.go
Normal file
@@ -0,0 +1,82 @@
|
||||
package playground
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"os"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"doormile/internal/testpg"
|
||||
)
|
||||
|
||||
// The booking executors against a real Postgres. Skipped unless
|
||||
// REGISTRY_TEST_DSN points at a THROWAWAY database (see
|
||||
// internal/ai/registry/store_integration_test.go for how to start one).
|
||||
//
|
||||
// The table is a minimal stand-in for pickupbookings WITH personal columns,
|
||||
// so the test proves the executors never select them.
|
||||
func TestBookingExecutorsSelectNoPersonalColumns(t *testing.T) {
|
||||
dsn := os.Getenv("REGISTRY_TEST_DSN")
|
||||
if dsn == "" {
|
||||
t.Skip("REGISTRY_TEST_DSN not set; skipping Postgres integration test")
|
||||
}
|
||||
db := testpg.Open(t, dsn, "aiplayground_test")
|
||||
for _, q := range []string{
|
||||
`DROP TABLE IF EXISTS pickupbookings`,
|
||||
`CREATE TABLE pickupbookings (
|
||||
bookingid serial PRIMARY KEY, bookingno text NOT NULL, tenantid int, status text,
|
||||
pickuppincode text, deliverypincode text, deliverycity text,
|
||||
pickuplatitude double precision, pickuplongitude double precision,
|
||||
deliverylatitude double precision, deliverylongitude double precision,
|
||||
assignedmileruserid int, routekm double precision, createdat timestamp,
|
||||
pickupaddress text, deliveryaddress text, notes text)`,
|
||||
`INSERT INTO pickupbookings (bookingno, status, pickuppincode, deliverypincode, deliverycity,
|
||||
pickuplatitude, pickuplongitude, deliverylatitude, deliverylongitude, createdat, pickupaddress, deliveryaddress, notes)
|
||||
VALUES ('DM-1', 'Created', '641001', '641002', 'Coimbatore', 11.01684, 76.95583, 11.0, 76.9, '2026-09-29 12:30', '12 MG Road', '4 Park St', 'call 9876543210'),
|
||||
('DM-2', 'Cancelled', '641001', '641003', 'Coimbatore', 11.0, 76.9, 11.0, 76.9, '2026-09-29 12:40', 'x', 'y', 'z')`,
|
||||
} {
|
||||
if err := db.Exec(q).Error; err != nil {
|
||||
t.Fatalf("setup: %v", err)
|
||||
}
|
||||
}
|
||||
execs := Executors(db, nil)
|
||||
if _, ok := execs["nearby_milers"]; ok {
|
||||
t.Fatal("nearby_milers offered without Redis")
|
||||
}
|
||||
|
||||
run := func(tool, input string) string {
|
||||
out, err := execs[tool](context.Background(), json.RawMessage(input))
|
||||
if err != nil {
|
||||
t.Fatalf("%s: %v", tool, err)
|
||||
}
|
||||
b, _ := json.Marshal(Redact(out))
|
||||
return string(b)
|
||||
}
|
||||
|
||||
got := run("get_booking_cache", `{"booking_id":1}`)
|
||||
for _, leaked := range []string{"MG Road", "Park St", "9876543210"} {
|
||||
if strings.Contains(got, leaked) {
|
||||
t.Fatalf("personal data %q in %s", leaked, got)
|
||||
}
|
||||
}
|
||||
for _, want := range []string{`"found":true`, `"pickuplat":11.02`, `"createdat_ist":"2026-09-29 12:30"`, `"status":"Created"`} {
|
||||
if !strings.Contains(got, want) {
|
||||
t.Fatalf("missing %s in %s", want, got)
|
||||
}
|
||||
}
|
||||
if !strings.Contains(run("get_booking_cache", `{"booking_id":99}`), `"found":false`) {
|
||||
t.Fatal("missing booking should be found:false")
|
||||
}
|
||||
if _, err := execs["get_booking_cache"](context.Background(), json.RawMessage(`{}`)); err == nil {
|
||||
t.Fatal("booking_id should be required")
|
||||
}
|
||||
|
||||
scan := run("scan_bookings", `{"status":"Cancelled"}`)
|
||||
if !strings.Contains(scan, `"count":1`) || !strings.Contains(scan, `"Cancelled":1`) {
|
||||
t.Fatalf("scan = %s", scan)
|
||||
}
|
||||
if all := run("scan_bookings", `{"limit":500}`); !strings.Contains(all, `"count":2`) {
|
||||
t.Fatalf("scan all = %s", all)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user