first commit
This commit is contained in:
@@ -0,0 +1,158 @@
|
||||
package permission
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
|
||||
"ai-agent/internal/db"
|
||||
)
|
||||
|
||||
type Policy string
|
||||
|
||||
const (
|
||||
PolicyAllow Policy = "allow"
|
||||
PolicyDeny Policy = "deny"
|
||||
PolicyAsk Policy = "ask"
|
||||
)
|
||||
|
||||
type Checker struct {
|
||||
store *db.Store
|
||||
cache map[string]Policy
|
||||
mu sync.RWMutex
|
||||
yolo bool
|
||||
}
|
||||
|
||||
func NewChecker(store *db.Store, yolo bool) *Checker {
|
||||
c := &Checker{
|
||||
store: store,
|
||||
cache: make(map[string]Policy),
|
||||
yolo: yolo,
|
||||
}
|
||||
if store != nil {
|
||||
c.loadFromDB()
|
||||
}
|
||||
return c
|
||||
}
|
||||
|
||||
func (c *Checker) Check(toolName string) Policy {
|
||||
if c.yolo {
|
||||
return PolicyAllow
|
||||
}
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
if p, ok := c.cache[toolName]; ok {
|
||||
return p
|
||||
}
|
||||
return PolicyAsk
|
||||
}
|
||||
|
||||
func (c *Checker) SetPolicy(toolName string, policy Policy) {
|
||||
c.mu.Lock()
|
||||
c.cache[toolName] = policy
|
||||
c.mu.Unlock()
|
||||
if c.store != nil {
|
||||
c.store.UpsertToolPermission(context.Background(), db.UpsertToolPermissionParams{
|
||||
ToolName: toolName,
|
||||
Policy: string(policy),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Checker) IsYolo() bool {
|
||||
return c.yolo
|
||||
}
|
||||
|
||||
func (c *Checker) AllPolicies() map[string]Policy {
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
result := make(map[string]Policy, len(c.cache))
|
||||
for k, v := range c.cache {
|
||||
result[k] = v
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func (c *Checker) Reset() {
|
||||
c.mu.Lock()
|
||||
c.cache = make(map[string]Policy)
|
||||
c.mu.Unlock()
|
||||
if c.store != nil {
|
||||
c.store.ResetToolPermissions(context.Background())
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Checker) loadFromDB() {
|
||||
perms, err := c.store.ListToolPermissions(context.Background())
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
for _, p := range perms {
|
||||
switch Policy(p.Policy) {
|
||||
case PolicyAllow, PolicyDeny, PolicyAsk:
|
||||
c.cache[p.ToolName] = Policy(p.Policy)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
type ApprovalRequest struct {
|
||||
ToolName string
|
||||
Args map[string]any
|
||||
Response chan ApprovalResponse
|
||||
}
|
||||
|
||||
type ApprovalResponse struct {
|
||||
Allowed bool
|
||||
Always bool
|
||||
}
|
||||
|
||||
func RequestApproval(toolName string, args map[string]any, callback func(ApprovalRequest)) (bool, bool) {
|
||||
if callback == nil {
|
||||
return true, false
|
||||
}
|
||||
ch := make(chan ApprovalResponse, 1)
|
||||
callback(ApprovalRequest{
|
||||
ToolName: toolName,
|
||||
Args: args,
|
||||
Response: ch,
|
||||
})
|
||||
resp := <-ch
|
||||
return resp.Allowed, resp.Always
|
||||
}
|
||||
|
||||
type CheckResult int
|
||||
|
||||
const (
|
||||
CheckAllow CheckResult = iota
|
||||
CheckDeny
|
||||
CheckAsk
|
||||
)
|
||||
|
||||
func (c *Checker) ToCheckResult(toolName string) CheckResult {
|
||||
if c == nil || c.yolo {
|
||||
return CheckAllow
|
||||
}
|
||||
switch c.Check(toolName) {
|
||||
case PolicyAllow:
|
||||
return CheckAllow
|
||||
case PolicyDeny:
|
||||
return CheckDeny
|
||||
default:
|
||||
return CheckAsk
|
||||
}
|
||||
}
|
||||
|
||||
func NilSafe(store *db.Store, yolo bool) *Checker {
|
||||
return NewChecker(store, yolo)
|
||||
}
|
||||
|
||||
var AlwaysAllow = func(_ ApprovalRequest) {}
|
||||
|
||||
type ErrDenied struct {
|
||||
ToolName string
|
||||
}
|
||||
|
||||
func (e *ErrDenied) Error() string {
|
||||
return "tool call denied by permission policy: " + e.ToolName
|
||||
}
|
||||
@@ -0,0 +1,98 @@
|
||||
package permission
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"ai-agent/internal/db"
|
||||
)
|
||||
|
||||
func TestChecker_DefaultPolicy(t *testing.T) {
|
||||
c := NewChecker(nil, false)
|
||||
if got := c.Check("some_tool"); got != PolicyAsk {
|
||||
t.Errorf("Check() = %q, want %q", got, PolicyAsk)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChecker_Yolo(t *testing.T) {
|
||||
c := NewChecker(nil, true)
|
||||
if got := c.Check("any_tool"); got != PolicyAllow {
|
||||
t.Errorf("Check() = %q, want %q", got, PolicyAllow)
|
||||
}
|
||||
if !c.IsYolo() {
|
||||
t.Error("expected IsYolo() = true")
|
||||
}
|
||||
}
|
||||
|
||||
func TestChecker_SetPolicy(t *testing.T) {
|
||||
c := NewChecker(nil, false)
|
||||
c.SetPolicy("bash", PolicyAllow)
|
||||
if got := c.Check("bash"); got != PolicyAllow {
|
||||
t.Errorf("Check() = %q, want %q", got, PolicyAllow)
|
||||
}
|
||||
c.SetPolicy("bash", PolicyDeny)
|
||||
if got := c.Check("bash"); got != PolicyDeny {
|
||||
t.Errorf("Check() = %q, want %q", got, PolicyDeny)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChecker_WithDB(t *testing.T) {
|
||||
store, err := db.OpenPath(filepath.Join(t.TempDir(), "test.db"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer store.Close()
|
||||
c := NewChecker(store, false)
|
||||
c.SetPolicy("file_write", PolicyAllow)
|
||||
c2 := NewChecker(store, false)
|
||||
if got := c2.Check("file_write"); got != PolicyAllow {
|
||||
t.Errorf("persisted Check() = %q, want %q", got, PolicyAllow)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChecker_Reset(t *testing.T) {
|
||||
c := NewChecker(nil, false)
|
||||
c.SetPolicy("tool1", PolicyAllow)
|
||||
c.SetPolicy("tool2", PolicyDeny)
|
||||
c.Reset()
|
||||
if got := c.Check("tool1"); got != PolicyAsk {
|
||||
t.Errorf("after reset Check() = %q, want %q", got, PolicyAsk)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChecker_AllPolicies(t *testing.T) {
|
||||
c := NewChecker(nil, false)
|
||||
c.SetPolicy("a", PolicyAllow)
|
||||
c.SetPolicy("b", PolicyDeny)
|
||||
|
||||
policies := c.AllPolicies()
|
||||
if len(policies) != 2 {
|
||||
t.Errorf("AllPolicies() len = %d, want 2", len(policies))
|
||||
}
|
||||
if policies["a"] != PolicyAllow {
|
||||
t.Errorf("policies[a] = %q, want %q", policies["a"], PolicyAllow)
|
||||
}
|
||||
}
|
||||
|
||||
func TestToCheckResult(t *testing.T) {
|
||||
c := NewChecker(nil, false)
|
||||
c.SetPolicy("allowed", PolicyAllow)
|
||||
c.SetPolicy("denied", PolicyDeny)
|
||||
|
||||
if c.ToCheckResult("allowed") != CheckAllow {
|
||||
t.Error("expected CheckAllow for allowed tool")
|
||||
}
|
||||
if c.ToCheckResult("denied") != CheckDeny {
|
||||
t.Error("expected CheckDeny for denied tool")
|
||||
}
|
||||
if c.ToCheckResult("unknown") != CheckAsk {
|
||||
t.Error("expected CheckAsk for unknown tool")
|
||||
}
|
||||
}
|
||||
|
||||
func TestToCheckResult_Nil(t *testing.T) {
|
||||
var c *Checker
|
||||
if c.ToCheckResult("anything") != CheckAllow {
|
||||
t.Error("nil checker should return CheckAllow")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user