feat(coderd): add Agent Firewall correlation columns to aibridge_interceptions (#24817)

Add `agent_firewall_session_id` (UUID NULL) and
`agent_firewall_sequence_number` (INT NULL) to `aibridge_interceptions`
with a partial index on `agent_firewall_session_id`. No FK to
`boundary_sessions` (soft reference, resolved at query time).
`RecordInterception` reads the new fields from the proto request (merged
in #25884) via `parseOptionalUUID` / `parseOptionalInt32` helpers.

> This PR was authored by Coder Agents.
This commit is contained in:
Sas Swart
2026-06-15 12:34:16 +02:00
committed by GitHub
parent bfe64d6355
commit e1c7e61eb9
11 changed files with 336 additions and 52 deletions
+45 -14
View File
@@ -180,21 +180,29 @@ func (s *Server) RecordInterception(ctx context.Context, in *proto.RecordInterce
providerName = in.Provider
}
agentFirewallSessionID, err := parseOptionalUUID(in.AgentFirewallSessionId)
if err != nil {
s.logger.Warn(ctx, "invalid agent firewall session ID in interception request",
slog.F("agent_firewall_session_id", in.GetAgentFirewallSessionId()), slog.Error(err))
}
_, err = s.store.InsertAIBridgeInterception(ctx, database.InsertAIBridgeInterceptionParams{
ID: intcID,
APIKeyID: sql.NullString{String: in.ApiKeyId, Valid: true},
Client: sql.NullString{String: in.Client, Valid: in.Client != ""},
ClientSessionID: sql.NullString{String: in.GetClientSessionId(), Valid: in.GetClientSessionId() != ""},
InitiatorID: initID,
Provider: in.Provider,
ProviderName: providerName,
Model: in.Model,
Metadata: out,
StartedAt: in.StartedAt.AsTime(),
ThreadParentInterceptionID: uuid.NullUUID{UUID: parentID, Valid: parentID != uuid.Nil},
ThreadRootInterceptionID: uuid.NullUUID{UUID: rootID, Valid: rootID != uuid.Nil},
CredentialKind: credentialKindOrDefault(in.CredentialKind),
CredentialHint: in.CredentialHint,
ID: intcID,
APIKeyID: sql.NullString{String: in.ApiKeyId, Valid: true},
Client: sql.NullString{String: in.Client, Valid: in.Client != ""},
ClientSessionID: sql.NullString{String: in.GetClientSessionId(), Valid: in.GetClientSessionId() != ""},
InitiatorID: initID,
Provider: in.Provider,
ProviderName: providerName,
Model: in.Model,
Metadata: out,
StartedAt: in.StartedAt.AsTime(),
ThreadParentInterceptionID: uuid.NullUUID{UUID: parentID, Valid: parentID != uuid.Nil},
ThreadRootInterceptionID: uuid.NullUUID{UUID: rootID, Valid: rootID != uuid.Nil},
CredentialKind: credentialKindOrDefault(in.CredentialKind),
CredentialHint: in.CredentialHint,
AgentFirewallSessionID: agentFirewallSessionID,
AgentFirewallSequenceNumber: parseOptionalInt32(in.AgentFirewallSequenceNumber),
})
if err != nil {
return nil, xerrors.Errorf("start interception: %w", err)
@@ -688,3 +696,26 @@ func metadataToMap(in map[string]*anypb.Any) map[string]any {
}
return meta
}
// parseOptionalUUID converts an optional proto string to uuid.NullUUID.
// Returns a zero NullUUID if s is nil. If s is non-nil but not a valid UUID, it
// returns a zero NullUUID along with the parse error so the caller can decide
// how to surface it.
func parseOptionalUUID(s *string) (uuid.NullUUID, error) {
if s == nil {
return uuid.NullUUID{}, nil
}
id, err := uuid.Parse(*s)
if err != nil {
return uuid.NullUUID{}, err
}
return uuid.NullUUID{UUID: id, Valid: true}, nil
}
// parseOptionalInt32 converts an optional proto int32 to sql.NullInt32.
func parseOptionalInt32(n *int32) sql.NullInt32 {
if n == nil {
return sql.NullInt32{}
}
return sql.NullInt32{Int32: *n, Valid: true}
}
@@ -688,6 +688,128 @@ func TestRecordInterception(t *testing.T) {
}, nil)
},
},
{
name: "valid interception with agent firewall correlation",
request: &proto.RecordInterceptionRequest{
Id: uuid.NewString(),
ApiKeyId: uuid.NewString(),
InitiatorId: uuid.NewString(),
Provider: "anthropic",
Model: "claude-4-opus",
Metadata: metadataProto,
StartedAt: timestamppb.Now(),
AgentFirewallSessionId: ptr.Ref(uuid.NewString()),
AgentFirewallSequenceNumber: ptr.Ref(int32(42)),
},
setupMocks: func(t *testing.T, db *dbmock.MockStore, req *proto.RecordInterceptionRequest) {
interceptionID, err := uuid.Parse(req.GetId())
assert.NoError(t, err, "parse interception UUID")
initiatorID, err := uuid.Parse(req.GetInitiatorId())
assert.NoError(t, err, "parse interception initiator UUID")
agentFirewallSessionID, err := uuid.Parse(req.GetAgentFirewallSessionId())
assert.NoError(t, err, "parse agent firewall session UUID")
db.EXPECT().InsertAIBridgeInterception(gomock.Any(), database.InsertAIBridgeInterceptionParams{
ID: interceptionID,
APIKeyID: sql.NullString{String: req.ApiKeyId, Valid: true},
InitiatorID: initiatorID,
Provider: req.GetProvider(),
ProviderName: req.GetProvider(),
Model: req.GetModel(),
Metadata: json.RawMessage(metadataJSON),
StartedAt: req.StartedAt.AsTime().UTC(),
CredentialKind: database.CredentialKindCentralized,
AgentFirewallSessionID: uuid.NullUUID{UUID: agentFirewallSessionID, Valid: true},
AgentFirewallSequenceNumber: sql.NullInt32{Int32: 42, Valid: true},
}).Return(database.AIBridgeInterception{
ID: interceptionID,
InitiatorID: initiatorID,
Provider: req.GetProvider(),
Model: req.GetModel(),
StartedAt: req.StartedAt.AsTime().UTC(),
}, nil)
},
},
{
name: "absent agent firewall fields treated as null",
request: &proto.RecordInterceptionRequest{
Id: uuid.NewString(),
ApiKeyId: uuid.NewString(),
InitiatorId: uuid.NewString(),
Provider: "anthropic",
Model: "claude-4-opus",
Metadata: metadataProto,
StartedAt: timestamppb.Now(),
},
setupMocks: func(t *testing.T, db *dbmock.MockStore, req *proto.RecordInterceptionRequest) {
interceptionID, err := uuid.Parse(req.GetId())
assert.NoError(t, err, "parse interception UUID")
initiatorID, err := uuid.Parse(req.GetInitiatorId())
assert.NoError(t, err, "parse interception initiator UUID")
db.EXPECT().InsertAIBridgeInterception(gomock.Any(), database.InsertAIBridgeInterceptionParams{
ID: interceptionID,
APIKeyID: sql.NullString{String: req.ApiKeyId, Valid: true},
InitiatorID: initiatorID,
Provider: req.GetProvider(),
ProviderName: req.GetProvider(),
Model: req.GetModel(),
Metadata: json.RawMessage(metadataJSON),
StartedAt: req.StartedAt.AsTime().UTC(),
CredentialKind: database.CredentialKindCentralized,
AgentFirewallSessionID: uuid.NullUUID{},
AgentFirewallSequenceNumber: sql.NullInt32{},
}).Return(database.AIBridgeInterception{
ID: interceptionID,
InitiatorID: initiatorID,
Provider: req.GetProvider(),
Model: req.GetModel(),
StartedAt: req.StartedAt.AsTime().UTC(),
}, nil)
},
},
{
name: "invalid agent firewall session ID treated as null",
request: &proto.RecordInterceptionRequest{
Id: uuid.NewString(),
ApiKeyId: uuid.NewString(),
InitiatorId: uuid.NewString(),
Provider: "anthropic",
Model: "claude-4-opus",
Metadata: metadataProto,
StartedAt: timestamppb.Now(),
AgentFirewallSessionId: ptr.Ref("not-a-uuid"),
AgentFirewallSequenceNumber: ptr.Ref(int32(7)),
},
setupMocks: func(t *testing.T, db *dbmock.MockStore, req *proto.RecordInterceptionRequest) {
interceptionID, err := uuid.Parse(req.GetId())
assert.NoError(t, err, "parse interception UUID")
initiatorID, err := uuid.Parse(req.GetInitiatorId())
assert.NoError(t, err, "parse interception initiator UUID")
// Malformed agent firewall session ID is stored as null
// (and logged) rather than failing the interception.
db.EXPECT().InsertAIBridgeInterception(gomock.Any(), database.InsertAIBridgeInterceptionParams{
ID: interceptionID,
APIKeyID: sql.NullString{String: req.ApiKeyId, Valid: true},
InitiatorID: initiatorID,
Provider: req.GetProvider(),
ProviderName: req.GetProvider(),
Model: req.GetModel(),
Metadata: json.RawMessage(metadataJSON),
StartedAt: req.StartedAt.AsTime().UTC(),
CredentialKind: database.CredentialKindCentralized,
AgentFirewallSessionID: uuid.NullUUID{},
AgentFirewallSequenceNumber: sql.NullInt32{Int32: 7, Valid: true},
}).Return(database.AIBridgeInterception{
ID: interceptionID,
InitiatorID: initiatorID,
Provider: req.GetProvider(),
Model: req.GetModel(),
StartedAt: req.StartedAt.AsTime().UTC(),
}, nil)
},
},
{
name: "invalid interception ID",
request: &proto.RecordInterceptionRequest{