fix: scope project skills by tenant
This commit is contained in:
@@ -6,6 +6,7 @@ import (
|
||||
"strings"
|
||||
|
||||
"git.warky.dev/wdevs/amcs/internal/generatedmodels"
|
||||
"git.warky.dev/wdevs/amcs/internal/tenancy"
|
||||
ext "git.warky.dev/wdevs/amcs/internal/types"
|
||||
)
|
||||
|
||||
@@ -307,21 +308,35 @@ func (db *DB) GetGuardrail(ctx context.Context, id int64) (ext.AgentGuardrail, e
|
||||
// Project Skills
|
||||
|
||||
func (db *DB) AddProjectSkill(ctx context.Context, projectID, skillID int64, override bool) error {
|
||||
_, err := db.pool.Exec(ctx, `
|
||||
args := []any{projectID, skillID, override}
|
||||
tenantWhere := projectSkillTenantWhere(ctx, &args, "p", "s")
|
||||
tag, err := db.pool.Exec(ctx, `
|
||||
insert into project_skills (project_id, skill_id, override)
|
||||
values ($1, $2, $3)
|
||||
select $1, $2, $3
|
||||
from projects p
|
||||
join agent_skills s on s.id = $2
|
||||
where p.id = $1`+tenantWhere+`
|
||||
on conflict (project_id, skill_id) do update set override = excluded.override
|
||||
`, projectID, skillID, override)
|
||||
`, args...)
|
||||
if err != nil {
|
||||
return fmt.Errorf("add project skill: %w", err)
|
||||
}
|
||||
if tag.RowsAffected() == 0 {
|
||||
return fmt.Errorf("project or skill not found")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (db *DB) RemoveProjectSkill(ctx context.Context, projectID, skillID int64) error {
|
||||
args := []any{projectID, skillID}
|
||||
tenantWhere := projectSkillTenantWhere(ctx, &args, "p", "s")
|
||||
tag, err := db.pool.Exec(ctx, `
|
||||
delete from project_skills where project_id = $1 and skill_id = $2
|
||||
`, projectID, skillID)
|
||||
delete from project_skills ps
|
||||
using projects p, agent_skills s
|
||||
where ps.project_id = $1
|
||||
and ps.skill_id = $2
|
||||
and p.id = ps.project_id
|
||||
and s.id = ps.skill_id`+tenantWhere, args...)
|
||||
if err != nil {
|
||||
return fmt.Errorf("remove project skill: %w", err)
|
||||
}
|
||||
@@ -332,15 +347,18 @@ func (db *DB) RemoveProjectSkill(ctx context.Context, projectID, skillID int64)
|
||||
}
|
||||
|
||||
func (db *DB) ListProjectSkills(ctx context.Context, projectID int64) ([]ext.AgentSkill, error) {
|
||||
args := []any{projectID}
|
||||
tenantWhere := projectSkillTenantWhere(ctx, &args, "p", "s")
|
||||
rows, err := db.pool.Query(ctx, `
|
||||
select s.id, s.name, s.description, s.content, s.tags::text[],
|
||||
s.language_tags::text[], s.library_tags::text[], s.framework_tags::text[],
|
||||
s.domain_tags::text[], s.created_at, s.updated_at, ps.override
|
||||
from agent_skills s
|
||||
join project_skills ps on ps.skill_id = s.id
|
||||
where ps.project_id = $1
|
||||
join projects p on p.id = ps.project_id
|
||||
where ps.project_id = $1`+tenantWhere+`
|
||||
order by s.name
|
||||
`, projectID)
|
||||
`, args...)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list project skills: %w", err)
|
||||
}
|
||||
@@ -360,6 +378,15 @@ func (db *DB) ListProjectSkills(ctx context.Context, projectID int64) ([]ext.Age
|
||||
return skills, rows.Err()
|
||||
}
|
||||
|
||||
func projectSkillTenantWhere(ctx context.Context, args *[]any, projectAlias, skillAlias string) string {
|
||||
if key, ok := tenancy.KeyFromContext(ctx); ok {
|
||||
*args = append(*args, key)
|
||||
placeholder := fmt.Sprintf("$%d", len(*args))
|
||||
return fmt.Sprintf(" and %s.tenant_id = %s and %s.tenant_id = %s", projectAlias, placeholder, skillAlias, placeholder)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// Project Guardrails
|
||||
|
||||
func (db *DB) AddProjectGuardrail(ctx context.Context, projectID, guardrailID int64) error {
|
||||
|
||||
@@ -0,0 +1,32 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/amcs/internal/tenancy"
|
||||
)
|
||||
|
||||
func TestProjectSkillTenantWhere(t *testing.T) {
|
||||
args := []any{int64(1), int64(2), true}
|
||||
where := projectSkillTenantWhere(tenancy.WithTenantKey(context.Background(), "tenant-a"), &args, "p", "s")
|
||||
|
||||
if where != " and p.tenant_id = $4 and s.tenant_id = $4" {
|
||||
t.Fatalf("where = %q", where)
|
||||
}
|
||||
if len(args) != 4 || args[3] != "tenant-a" {
|
||||
t.Fatalf("args = %#v", args)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProjectSkillTenantWhereWithoutTenant(t *testing.T) {
|
||||
args := []any{int64(1), int64(2)}
|
||||
where := projectSkillTenantWhere(context.Background(), &args, "p", "s")
|
||||
|
||||
if where != "" {
|
||||
t.Fatalf("where = %q, want empty", where)
|
||||
}
|
||||
if len(args) != 2 {
|
||||
t.Fatalf("args = %#v", args)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user