diff --git a/pkg/tools/obligation.go b/pkg/tools/obligation.go new file mode 100644 index 000000000..44867f426 --- /dev/null +++ b/pkg/tools/obligation.go @@ -0,0 +1,439 @@ +package tools + +import ( + "context" + "fmt" + "sort" + "strings" + "time" + + jsonv2 "github.com/go-json-experiment/json" + + "github.com/ZanzyTHEbar/dragonscale/pkg/ids" + "github.com/ZanzyTHEbar/dragonscale/pkg/memory" +) + +const obligationKVPrefix = "obligation:" + +type ObligationState string + +const ( + ObligationStateCreated ObligationState = "created" + ObligationStateScheduled ObligationState = "scheduled" + ObligationStateDue ObligationState = "due" + ObligationStateExecuted ObligationState = "executed" + ObligationStateVerified ObligationState = "verified" +) + +type ObligationEvidence struct { + At time.Time `json:"at"` + Source string `json:"source"` + Content string `json:"content"` +} + +type ObligationRecord struct { + ID string `json:"id"` + Title string `json:"title"` + Details string `json:"details,omitempty"` + State ObligationState `json:"state"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` + ScheduledAt time.Time `json:"scheduled_at,omitempty"` + DueAt time.Time `json:"due_at,omitempty"` + ExecutedAt time.Time `json:"executed_at,omitempty"` + VerifiedAt time.Time `json:"verified_at,omitempty"` + Evidence []ObligationEvidence `json:"evidence,omitempty"` +} + +type ObligationTool struct { + delegate memory.MemoryDelegate + agentID string +} + +func NewObligationTool(delegate memory.MemoryDelegate, agentID string) *ObligationTool { + return &ObligationTool{ + delegate: delegate, + agentID: agentID, + } +} + +func (t *ObligationTool) Name() string { + return "obligation" +} + +func (t *ObligationTool) Description() string { + return "Manage commitment/reminder obligations with lifecycle states and audit evidence." +} + +func (t *ObligationTool) Parameters() map[string]interface{} { + return map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "action": map[string]interface{}{ + "type": "string", + "description": "create|get|list|update_state|add_evidence", + }, + "obligation_id": map[string]interface{}{ + "type": "string", + "description": "Obligation ID for get/update_state/add_evidence.", + }, + "title": map[string]interface{}{ + "type": "string", + "description": "Title for create action.", + }, + "details": map[string]interface{}{ + "type": "string", + "description": "Optional details for create action.", + }, + "scheduled_at": map[string]interface{}{ + "type": "string", + "description": "Optional RFC3339 schedule time.", + }, + "due_at": map[string]interface{}{ + "type": "string", + "description": "Optional RFC3339 due time.", + }, + "state": map[string]interface{}{ + "type": "string", + "description": "Target lifecycle state for update_state.", + }, + "evidence": map[string]interface{}{ + "type": "string", + "description": "Evidence text for add_evidence.", + }, + "source": map[string]interface{}{ + "type": "string", + "description": "Optional evidence source label.", + }, + }, + "required": []string{"action"}, + } +} + +func (t *ObligationTool) Execute(ctx context.Context, args map[string]interface{}) *ToolResult { + if t.delegate == nil { + return ErrorResult("obligation store is not configured").WithError(fmt.Errorf("obligation delegate is nil")) + } + action, _ := args["action"].(string) + switch strings.ToLower(strings.TrimSpace(action)) { + case "create": + return t.create(ctx, args) + case "get": + return t.get(ctx, args) + case "list": + return t.list(ctx) + case "update_state": + return t.updateState(ctx, args) + case "add_evidence": + return t.addEvidence(ctx, args) + default: + return ErrorResult("unknown obligation action").WithError(fmt.Errorf("unknown obligation action: %s", action)) + } +} + +func (t *ObligationTool) create(ctx context.Context, args map[string]interface{}) *ToolResult { + title, _ := args["title"].(string) + if strings.TrimSpace(title) == "" { + return ErrorResult("title is required for create").WithError(fmt.Errorf("title is required")) + } + now := time.Now().UTC() + rec := &ObligationRecord{ + ID: ids.New().String(), + Title: title, + Details: stringOr(args["details"]), + State: ObligationStateCreated, + CreatedAt: now, + UpdatedAt: now, + } + + if scheduledAtRaw := stringOr(args["scheduled_at"]); scheduledAtRaw != "" { + ts, err := time.Parse(time.RFC3339, scheduledAtRaw) + if err != nil { + return ErrorResult("scheduled_at must be RFC3339").WithError(err) + } + rec.ScheduledAt = ts.UTC() + rec.State = ObligationStateScheduled + } + if dueAtRaw := stringOr(args["due_at"]); dueAtRaw != "" { + ts, err := time.Parse(time.RFC3339, dueAtRaw) + if err != nil { + return ErrorResult("due_at must be RFC3339").WithError(err) + } + rec.DueAt = ts.UTC() + if rec.State == ObligationStateCreated { + rec.State = ObligationStateScheduled + } + } + + if err := t.saveRecord(ctx, rec); err != nil { + return ErrorResult("failed to persist obligation").WithError(err) + } + return obligationSuccess(rec) +} + +func (t *ObligationTool) get(ctx context.Context, args map[string]interface{}) *ToolResult { + id := stringOr(args["obligation_id"]) + if id == "" { + return ErrorResult("obligation_id is required").WithError(fmt.Errorf("obligation_id is required")) + } + rec, err := t.loadRecord(ctx, id) + if err != nil { + return ErrorResult("failed to load obligation").WithError(err) + } + return obligationSuccess(rec) +} + +func (t *ObligationTool) list(ctx context.Context) *ToolResult { + rows, err := t.delegate.ListKVByPrefix(ctx, t.agentID, obligationKVPrefix, 500) + if err != nil { + return ErrorResult("failed to list obligations").WithError(err) + } + result := make([]*ObligationRecord, 0, len(rows)) + for _, raw := range rows { + var rec ObligationRecord + if err := jsonv2.Unmarshal([]byte(raw), &rec); err != nil { + continue + } + result = append(result, &rec) + } + data, err := jsonv2.Marshal(map[string]interface{}{ + "count": len(result), + "obligations": result, + }) + if err != nil { + return ErrorResult("failed to serialize obligations").WithError(err) + } + return &ToolResult{ForLLM: string(data), ForUser: string(data), Silent: false, IsError: false} +} + +// CollectDueObligations returns obligations that are due at or before now. +// Scheduled obligations that become due are atomically transitioned to "due" +// with a scheduler evidence record. +func (t *ObligationTool) CollectDueObligations(ctx context.Context, now time.Time, source string) ([]*ObligationRecord, error) { + if t.delegate == nil { + return nil, fmt.Errorf("obligation delegate is nil") + } + + now = now.UTC() + source = strings.TrimSpace(source) + + rows, err := t.delegate.ListKVByPrefix(ctx, t.agentID, obligationKVPrefix, 500) + if err != nil { + return nil, fmt.Errorf("list obligations: %w", err) + } + + due := make([]*ObligationRecord, 0, len(rows)) + for _, raw := range rows { + var rec ObligationRecord + if err := jsonv2.Unmarshal([]byte(raw), &rec); err != nil { + continue + } + if !obligationIsDue(&rec, now) { + continue + } + + if rec.State == ObligationStateScheduled { + rec.State = ObligationStateDue + rec.UpdatedAt = now + if rec.DueAt.IsZero() { + rec.DueAt = now + } + if source != "" { + rec.Evidence = append(rec.Evidence, ObligationEvidence{ + At: now, + Source: source, + Content: "obligation became due via scheduler check", + }) + } + if err := t.saveRecord(ctx, &rec); err != nil { + return nil, fmt.Errorf("persist due obligation %s: %w", rec.ID, err) + } + } + + recCopy := rec + due = append(due, &recCopy) + } + + sort.SliceStable(due, func(i, j int) bool { + left := obligationSortTime(due[i]) + right := obligationSortTime(due[j]) + if left.Equal(right) { + return due[i].ID < due[j].ID + } + return left.Before(right) + }) + + return due, nil +} + +func (t *ObligationTool) updateState(ctx context.Context, args map[string]interface{}) *ToolResult { + id := stringOr(args["obligation_id"]) + if id == "" { + return ErrorResult("obligation_id is required").WithError(fmt.Errorf("obligation_id is required")) + } + next := ObligationState(stringOr(args["state"])) + if !isValidObligationState(next) { + return ErrorResult("invalid obligation state").WithError(fmt.Errorf("invalid obligation state: %s", next)) + } + rec, err := t.loadRecord(ctx, id) + if err != nil { + return ErrorResult("failed to load obligation").WithError(err) + } + if err := validateObligationTransition(rec.State, next, len(rec.Evidence)); err != nil { + return ErrorResult(err.Error()).WithError(err) + } + + now := time.Now().UTC() + rec.State = next + rec.UpdatedAt = now + switch next { + case ObligationStateDue: + if rec.DueAt.IsZero() { + rec.DueAt = now + } + case ObligationStateExecuted: + rec.ExecutedAt = now + case ObligationStateVerified: + rec.VerifiedAt = now + } + if err := t.saveRecord(ctx, rec); err != nil { + return ErrorResult("failed to persist obligation state").WithError(err) + } + return obligationSuccess(rec) +} + +func (t *ObligationTool) addEvidence(ctx context.Context, args map[string]interface{}) *ToolResult { + id := stringOr(args["obligation_id"]) + if id == "" { + return ErrorResult("obligation_id is required").WithError(fmt.Errorf("obligation_id is required")) + } + evidence := stringOr(args["evidence"]) + if strings.TrimSpace(evidence) == "" { + return ErrorResult("evidence is required").WithError(fmt.Errorf("evidence is required")) + } + rec, err := t.loadRecord(ctx, id) + if err != nil { + return ErrorResult("failed to load obligation").WithError(err) + } + rec.Evidence = append(rec.Evidence, ObligationEvidence{ + At: time.Now().UTC(), + Source: stringOr(args["source"]), + Content: evidence, + }) + rec.UpdatedAt = time.Now().UTC() + if err := t.saveRecord(ctx, rec); err != nil { + return ErrorResult("failed to persist obligation evidence").WithError(err) + } + return obligationSuccess(rec) +} + +func (t *ObligationTool) loadRecord(ctx context.Context, id string) (*ObligationRecord, error) { + raw, err := t.delegate.GetKV(ctx, t.agentID, obligationKVPrefix+id) + if err != nil { + return nil, err + } + if strings.TrimSpace(raw) == "" { + return nil, fmt.Errorf("obligation %s not found", id) + } + var rec ObligationRecord + if err := jsonv2.Unmarshal([]byte(raw), &rec); err != nil { + return nil, err + } + return &rec, nil +} + +func (t *ObligationTool) saveRecord(ctx context.Context, rec *ObligationRecord) error { + data, err := jsonv2.Marshal(rec) + if err != nil { + return err + } + return t.delegate.UpsertKV(ctx, t.agentID, obligationKVPrefix+rec.ID, string(data)) +} + +func validateObligationTransition(current, next ObligationState, evidenceCount int) error { + if current == next { + return nil + } + allowed := map[ObligationState]map[ObligationState]bool{ + ObligationStateCreated: { + ObligationStateScheduled: true, + ObligationStateDue: true, + }, + ObligationStateScheduled: { + ObligationStateDue: true, + ObligationStateExecuted: true, + }, + ObligationStateDue: { + ObligationStateExecuted: true, + }, + ObligationStateExecuted: { + ObligationStateVerified: true, + }, + ObligationStateVerified: {}, + } + if !allowed[current][next] { + return fmt.Errorf("invalid obligation transition: %s -> %s", current, next) + } + if next == ObligationStateVerified && evidenceCount == 0 { + return fmt.Errorf("verified state requires evidence") + } + return nil +} + +func obligationIsDue(rec *ObligationRecord, now time.Time) bool { + switch rec.State { + case ObligationStateScheduled: + if !rec.DueAt.IsZero() { + return !rec.DueAt.After(now) + } + if !rec.ScheduledAt.IsZero() { + return !rec.ScheduledAt.After(now) + } + return false + case ObligationStateDue: + return true + default: + return false + } +} + +func obligationSortTime(rec *ObligationRecord) time.Time { + if rec == nil { + return time.Time{} + } + if !rec.DueAt.IsZero() { + return rec.DueAt + } + if !rec.ScheduledAt.IsZero() { + return rec.ScheduledAt + } + if !rec.UpdatedAt.IsZero() { + return rec.UpdatedAt + } + return rec.CreatedAt +} + +func isValidObligationState(s ObligationState) bool { + switch s { + case ObligationStateCreated, ObligationStateScheduled, ObligationStateDue, ObligationStateExecuted, ObligationStateVerified: + return true + default: + return false + } +} + +func obligationSuccess(rec *ObligationRecord) *ToolResult { + data, _ := jsonv2.Marshal(rec) + return &ToolResult{ + ForLLM: string(data), + ForUser: string(data), + Silent: false, + IsError: false, + Async: false, + } +} + +func stringOr(v interface{}) string { + s, _ := v.(string) + return s +} diff --git a/pkg/tools/obligation_test.go b/pkg/tools/obligation_test.go new file mode 100644 index 000000000..4a1655a4a --- /dev/null +++ b/pkg/tools/obligation_test.go @@ -0,0 +1,183 @@ +package tools + +import ( + "context" + "testing" + "time" + + jsonv2 "github.com/go-json-experiment/json" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/ZanzyTHEbar/dragonscale/pkg/memory/delegate" +) + +func TestObligationTool_CreateAndList(t *testing.T) { + ctx := context.Background() + del, err := delegate.NewLibSQLInMemory() + require.NoError(t, err) + require.NoError(t, del.Init(ctx)) + defer del.Close() + + tool := NewObligationTool(del, "test-agent") + dueAt := time.Now().UTC().Add(4 * time.Hour).Format(time.RFC3339) + create := tool.Execute(ctx, map[string]interface{}{ + "action": "create", + "title": "send weekly update", + "due_at": dueAt, + }) + require.NotNil(t, create) + require.False(t, create.IsError, create.ForLLM) + + var rec ObligationRecord + require.NoError(t, jsonv2.Unmarshal([]byte(create.ForLLM), &rec)) + require.NotEmpty(t, rec.ID) + assert.Equal(t, ObligationStateScheduled, rec.State) + + list := tool.Execute(ctx, map[string]interface{}{"action": "list"}) + require.NotNil(t, list) + require.False(t, list.IsError, list.ForLLM) + + var payload struct { + Count int `json:"count"` + } + require.NoError(t, jsonv2.Unmarshal([]byte(list.ForLLM), &payload)) + assert.GreaterOrEqual(t, payload.Count, 1) +} + +func TestObligationTool_StateMachineAndEvidence(t *testing.T) { + ctx := context.Background() + 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", + "title": "follow up with candidate", + }) + require.False(t, create.IsError, create.ForLLM) + + var rec ObligationRecord + require.NoError(t, jsonv2.Unmarshal([]byte(create.ForLLM), &rec)) + + invalid := tool.Execute(ctx, map[string]interface{}{ + "action": "update_state", + "obligation_id": rec.ID, + "state": "verified", + }) + require.True(t, invalid.IsError) + assert.Contains(t, invalid.ForLLM, "invalid obligation transition") + + toDue := tool.Execute(ctx, map[string]interface{}{ + "action": "update_state", + "obligation_id": rec.ID, + "state": "due", + }) + require.False(t, toDue.IsError, toDue.ForLLM) + + toExecuted := tool.Execute(ctx, map[string]interface{}{ + "action": "update_state", + "obligation_id": rec.ID, + "state": "executed", + }) + require.False(t, toExecuted.IsError, toExecuted.ForLLM) + + missingEvidence := tool.Execute(ctx, map[string]interface{}{ + "action": "update_state", + "obligation_id": rec.ID, + "state": "verified", + }) + require.True(t, missingEvidence.IsError) + assert.Contains(t, missingEvidence.ForLLM, "requires evidence") + + withEvidence := tool.Execute(ctx, map[string]interface{}{ + "action": "add_evidence", + "obligation_id": rec.ID, + "evidence": "sent confirmation email", + "source": "email", + }) + require.False(t, withEvidence.IsError, withEvidence.ForLLM) + + toVerified := tool.Execute(ctx, map[string]interface{}{ + "action": "update_state", + "obligation_id": rec.ID, + "state": "verified", + }) + require.False(t, toVerified.IsError, toVerified.ForLLM) + + var verified ObligationRecord + require.NoError(t, jsonv2.Unmarshal([]byte(toVerified.ForLLM), &verified)) + assert.Equal(t, ObligationStateVerified, verified.State) + assert.NotZero(t, verified.VerifiedAt) + require.Len(t, verified.Evidence, 1) +} + +func TestObligationTool_CollectDueObligations_TransitionsScheduledToDue(t *testing.T) { + ctx := context.Background() + 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", + "title": "send reminder", + "due_at": time.Now().UTC().Add(-5 * time.Minute).Format(time.RFC3339), + }) + require.False(t, create.IsError, create.ForLLM) + + var created ObligationRecord + require.NoError(t, jsonv2.Unmarshal([]byte(create.ForLLM), &created)) + require.Equal(t, ObligationStateScheduled, created.State) + + due, err := tool.CollectDueObligations(ctx, time.Now().UTC(), "heartbeat") + require.NoError(t, err) + require.Len(t, due, 1) + assert.Equal(t, created.ID, due[0].ID) + assert.Equal(t, ObligationStateDue, due[0].State) + require.NotEmpty(t, due[0].Evidence) + assert.Equal(t, "heartbeat", due[0].Evidence[0].Source) + + get := tool.Execute(ctx, map[string]interface{}{ + "action": "get", + "obligation_id": created.ID, + }) + require.False(t, get.IsError, get.ForLLM) + + var persisted ObligationRecord + require.NoError(t, jsonv2.Unmarshal([]byte(get.ForLLM), &persisted)) + assert.Equal(t, ObligationStateDue, persisted.State) + require.Len(t, persisted.Evidence, 1) +} + +func TestObligationTool_CollectDueObligations_DoesNotDuplicateDueEvidence(t *testing.T) { + ctx := context.Background() + 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", + "title": "follow up now", + "due_at": time.Now().UTC().Add(-1 * time.Minute).Format(time.RFC3339), + }) + require.False(t, create.IsError, create.ForLLM) + + var created ObligationRecord + require.NoError(t, jsonv2.Unmarshal([]byte(create.ForLLM), &created)) + + first, err := tool.CollectDueObligations(ctx, time.Now().UTC(), "heartbeat") + require.NoError(t, err) + require.Len(t, first, 1) + + second, err := tool.CollectDueObligations(ctx, time.Now().UTC().Add(1*time.Minute), "heartbeat") + require.NoError(t, err) + require.Len(t, second, 1) + require.Len(t, second[0].Evidence, 1, "due transition evidence should be appended only once") + assert.Equal(t, created.ID, second[0].ID) +}