mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
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:
@@ -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{
|
||||
|
||||
Reference in New Issue
Block a user