Skip to content

Commit fa543de

Browse files
authored
Scope an init script to the environment it targets (#103)
1 parent ad4b333 commit fa543de

3 files changed

Lines changed: 95 additions & 15 deletions

File tree

internal/server/environment_authorization_db_test.go

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -50,6 +50,9 @@ func seedEnvironment(ctx context.Context, t *testing.T, backing *store.Store, or
5050
RunnerID: &runnerID,
5151
Flavor: "small",
5252
Availability: store.EnvironmentAvailabilityPrivate,
53+
LLMMode: store.LLMModePlatform,
54+
// NOT NULL, and a nil slice encodes as NULL rather than '{}'.
55+
LLMAllowedModels: []string{},
5356
})
5457
if err != nil {
5558
t.Fatalf("seed environment: %v", err)
Lines changed: 86 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,86 @@
1+
package server
2+
3+
import (
4+
"context"
5+
"testing"
6+
7+
"github.com/agynio/agents/internal/store"
8+
"github.com/google/uuid"
9+
)
10+
11+
// An init script targets an agent, an mcp or an environment. Only the first two
12+
// were ever filtered on, so a list scoped to one environment returned every
13+
// script in the database -- and every mutation tried to resolve an agent a
14+
// environment-scoped row does not have, which rolled the write back.
15+
func seedEnvironmentInitScript(ctx context.Context, t *testing.T, backing *store.Store, organizationID, environmentID uuid.UUID, script string) uuid.UUID {
16+
t.Helper()
17+
created, err := backing.CreateInitScript(ctx, organizationID, store.InitScriptInput{
18+
Script: script,
19+
Description: script,
20+
EnvironmentID: &environmentID,
21+
})
22+
if err != nil {
23+
t.Fatalf("create init script: %v", err)
24+
}
25+
return created.Meta.ID
26+
}
27+
28+
func TestListInitScriptsIsScopedToItsEnvironment(t *testing.T) {
29+
ctx := context.Background()
30+
_, backing := environmentAuthorizationServer(ctx, t, &recordingAuthorizationWriter{})
31+
organizationID := uuid.New()
32+
first := seedEnvironment(ctx, t, backing, organizationID, "first")
33+
second := seedEnvironment(ctx, t, backing, organizationID, "second")
34+
seedEnvironmentInitScript(ctx, t, backing, organizationID, first, "echo first")
35+
36+
result, err := backing.ListInitScripts(ctx, store.InitScriptFilter{EnvironmentID: &first}, 100, nil)
37+
if err != nil {
38+
t.Fatalf("list first: %v", err)
39+
}
40+
if len(result.InitScripts) != 1 {
41+
t.Fatalf("expected the environment's own script, got %d", len(result.InitScripts))
42+
}
43+
44+
// The leak: an unfiltered query returns this one too.
45+
other, err := backing.ListInitScripts(ctx, store.InitScriptFilter{EnvironmentID: &second}, 100, nil)
46+
if err != nil {
47+
t.Fatalf("list second: %v", err)
48+
}
49+
if len(other.InitScripts) != 0 {
50+
t.Fatalf("expected no scripts for an environment that has none, got %d", len(other.InitScripts))
51+
}
52+
}
53+
54+
func TestDeleteEnvironmentInitScript(t *testing.T) {
55+
ctx := context.Background()
56+
_, backing := environmentAuthorizationServer(ctx, t, &recordingAuthorizationWriter{})
57+
organizationID := uuid.New()
58+
environmentID := seedEnvironment(ctx, t, backing, organizationID, "first")
59+
scriptID := seedEnvironmentInitScript(ctx, t, backing, organizationID, environmentID, "echo first")
60+
61+
// Resolving an agent for a row that has none used to fail inside the
62+
// transaction, so the delete rolled back and reported an internal error.
63+
if err := backing.DeleteInitScript(ctx, scriptID); err != nil {
64+
t.Fatalf("delete init script: %v", err)
65+
}
66+
result, err := backing.ListInitScripts(ctx, store.InitScriptFilter{EnvironmentID: &environmentID}, 100, nil)
67+
if err != nil {
68+
t.Fatalf("list after delete: %v", err)
69+
}
70+
if len(result.InitScripts) != 0 {
71+
t.Fatalf("expected the script to be gone, got %d", len(result.InitScripts))
72+
}
73+
}
74+
75+
func TestUpdateEnvironmentInitScript(t *testing.T) {
76+
ctx := context.Background()
77+
_, backing := environmentAuthorizationServer(ctx, t, &recordingAuthorizationWriter{})
78+
organizationID := uuid.New()
79+
environmentID := seedEnvironment(ctx, t, backing, organizationID, "first")
80+
scriptID := seedEnvironmentInitScript(ctx, t, backing, organizationID, environmentID, "echo first")
81+
82+
updated := "echo second"
83+
if _, err := backing.UpdateInitScript(ctx, scriptID, store.InitScriptUpdate{Script: &updated}); err != nil {
84+
t.Fatalf("update init script: %v", err)
85+
}
86+
}

internal/store/store.go

Lines changed: 6 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -1127,11 +1127,7 @@ func (s *Store) CreateInitScript(ctx context.Context, organizationID uuid.UUID,
11271127
}
11281128
return InitScript{}, err
11291129
}
1130-
agentID, err := resolveAgentID(ctx, tx, script.AgentID, script.McpID)
1131-
if err != nil {
1132-
return InitScript{}, err
1133-
}
1134-
if err := touchAgentUpdatedAt(ctx, tx, agentID); err != nil {
1130+
if err := touchTargetAgent(ctx, tx, script.AgentID, script.McpID, script.EnvironmentID); err != nil {
11351131
return InitScript{}, err
11361132
}
11371133
return script, nil
@@ -1175,11 +1171,7 @@ func (s *Store) UpdateInitScript(ctx context.Context, id uuid.UUID, update InitS
11751171
}
11761172
return InitScript{}, err
11771173
}
1178-
agentID, err := resolveAgentID(ctx, tx, script.AgentID, script.McpID)
1179-
if err != nil {
1180-
return InitScript{}, err
1181-
}
1182-
if err := touchAgentUpdatedAt(ctx, tx, agentID); err != nil {
1174+
if err := touchTargetAgent(ctx, tx, script.AgentID, script.McpID, script.EnvironmentID); err != nil {
11831175
return InitScript{}, err
11841176
}
11851177
return script, nil
@@ -1199,11 +1191,7 @@ func (s *Store) DeleteInitScript(ctx context.Context, id uuid.UUID) error {
11991191
}
12001192
return struct{}{}, err
12011193
}
1202-
agentID, err := resolveAgentID(ctx, tx, script.AgentID, script.McpID)
1203-
if err != nil {
1204-
return struct{}{}, err
1205-
}
1206-
if err := touchAgentUpdatedAt(ctx, tx, agentID); err != nil {
1194+
if err := touchTargetAgent(ctx, tx, script.AgentID, script.McpID, script.EnvironmentID); err != nil {
12071195
return struct{}{}, err
12081196
}
12091197
return struct{}{}, nil
@@ -1220,6 +1208,9 @@ func (s *Store) ListInitScripts(ctx context.Context, filter InitScriptFilter, pa
12201208
if filter.McpID != nil {
12211209
clauses, args = appendClause(clauses, args, "mcp_id = $%d", *filter.McpID)
12221210
}
1211+
if filter.EnvironmentID != nil {
1212+
clauses, args = appendClause(clauses, args, "environment_id = $%d", *filter.EnvironmentID)
1213+
}
12231214

12241215
scripts, nextCursor, err := listEntities(ctx, s.pool,
12251216
fmt.Sprintf("SELECT %s FROM init_scripts", initScriptColumns),

0 commit comments

Comments
 (0)