Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 6 additions & 1 deletion Makefile
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
.PHONY: all build run clean tidy test test-race test-dashboard test-e2e test-integration test-contract test-all lint lint-fix record-api swagger docs-openapi install-tools perf-check perf-bench infra image
.PHONY: all build run clean tidy test test-race test-dashboard test-e2e test-integration test-contract test-all lint lint-fix record-api swagger docs-openapi install-tools perf-check perf-bench infra image seed-demo-data

all: build

Expand Down Expand Up @@ -42,6 +42,11 @@ infra:
image:
docker compose --profile app up -d

# Seed rolling demo usage/audit data into SQLite.
# Usage: SQLITE_PATH=data/gomodel.db make seed-demo-data
seed-demo-data:
bash tools/seed-demo-data.sh

# Run unit tests only
test:
go test ./cmd/... ./internal/... ./config/... -v
Expand Down
20 changes: 17 additions & 3 deletions internal/auditlog/reader_postgresql.go
Original file line number Diff line number Diff line change
Expand Up @@ -9,12 +9,18 @@ import (

"github.com/goccy/go-json"

"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgxpool"
)

type postgreSQLQueryer interface {
Query(ctx context.Context, sql string, args ...any) (pgx.Rows, error)
QueryRow(ctx context.Context, sql string, args ...any) pgx.Row
}

// PostgreSQLReader implements Reader for PostgreSQL databases.
type PostgreSQLReader struct {
pool *pgxpool.Pool
pool postgreSQLQueryer
}

// NewPostgreSQLReader creates a new PostgreSQL audit log reader.
Expand Down Expand Up @@ -115,9 +121,10 @@ func (r *PostgreSQLReader) GetLogs(ctx context.Context, params LogQueryParams) (
var authKeyID *string
var authMethod *string
var userPath *string
var errorType *string

if err := rows.Scan(&e.ID, &e.Timestamp, &e.DurationNs, &e.RequestedModel, &e.ResolvedModel, &e.Provider, &providerName, &e.AliasUsed, &workflowVersionID, &cacheType, &e.StatusCode,
&e.RequestID, &authKeyID, &authMethod, &e.ClientIP, &e.Method, &e.Path, &userPath, &e.Stream, &e.ErrorType, &dataJSON); err != nil {
&e.RequestID, &authKeyID, &authMethod, &e.ClientIP, &e.Method, &e.Path, &userPath, &e.Stream, &errorType, &dataJSON); err != nil {
return nil, fmt.Errorf("failed to scan audit log row: %w", err)
}
if workflowVersionID != nil {
Expand All @@ -140,6 +147,9 @@ func (r *PostgreSQLReader) GetLogs(ctx context.Context, params LogQueryParams) (
if userPath != nil {
e.UserPath = *userPath
}
if errorType != nil {
e.ErrorType = *errorType
}

if dataJSON != nil && *dataJSON != "" {
var data LogData
Expand Down Expand Up @@ -257,9 +267,10 @@ func scanPostgreSQLLogEntry(rows interface {
var authKeyID *string
var authMethod *string
var userPath *string
var errorType *string

if err := rows.Scan(&e.ID, &e.Timestamp, &e.DurationNs, &e.RequestedModel, &e.ResolvedModel, &e.Provider, &providerName, &e.AliasUsed, &workflowVersionID, &cacheType, &e.StatusCode,
&e.RequestID, &authKeyID, &authMethod, &e.ClientIP, &e.Method, &e.Path, &userPath, &e.Stream, &e.ErrorType, &dataJSON); err != nil {
&e.RequestID, &authKeyID, &authMethod, &e.ClientIP, &e.Method, &e.Path, &userPath, &e.Stream, &errorType, &dataJSON); err != nil {
return nil, fmt.Errorf("failed to scan audit log row: %w", err)
}
if workflowVersionID != nil {
Expand All @@ -282,6 +293,9 @@ func scanPostgreSQLLogEntry(rows interface {
if userPath != nil {
e.UserPath = *userPath
}
if errorType != nil {
e.ErrorType = *errorType
}

if dataJSON != nil && *dataJSON != "" {
var data LogData
Expand Down
202 changes: 202 additions & 0 deletions internal/auditlog/reader_postgresql_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,202 @@
package auditlog

import (
"context"
"fmt"
"reflect"
"strings"
"testing"
"time"

"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgconn"
)

type fakePostgreSQLRow struct {
values []any
}

func (r fakePostgreSQLRow) Scan(dest ...any) error {
if len(dest) != len(r.values) {
return fmt.Errorf("scan destination count = %d, want %d", len(dest), len(r.values))
}
for i, value := range r.values {
target := reflect.ValueOf(dest[i])
if target.Kind() != reflect.Pointer || target.IsNil() {
return fmt.Errorf("scan destination %d is not a non-nil pointer", i)
}
elem := target.Elem()
if value == nil {
elem.Set(reflect.Zero(elem.Type()))
continue
}
if elem.Kind() == reflect.Pointer {
pointerValue := reflect.New(elem.Type().Elem())
if err := assignScannedValue(pointerValue.Elem(), value); err != nil {
return fmt.Errorf("scan destination %d: %w", i, err)
}
elem.Set(pointerValue)
continue
}
if err := assignScannedValue(elem, value); err != nil {
return fmt.Errorf("scan destination %d: %w", i, err)
}
}
return nil
}

func assignScannedValue(target reflect.Value, value any) error {
source := reflect.ValueOf(value)
if source.Type().AssignableTo(target.Type()) {
target.Set(source)
return nil
}
if source.Type().ConvertibleTo(target.Type()) {
target.Set(source.Convert(target.Type()))
return nil
}
return fmt.Errorf("cannot assign %T to %s", value, target.Type())
}

type fakePostgreSQLQueryer struct {
count int
rows pgx.Rows
}

func (q fakePostgreSQLQueryer) QueryRow(_ context.Context, _ string, _ ...any) pgx.Row {
return fakePostgreSQLRow{values: []any{q.count}}
}

func (q fakePostgreSQLQueryer) Query(_ context.Context, sql string, _ ...any) (pgx.Rows, error) {
if !strings.Contains(sql, "FROM audit_logs") {
return nil, fmt.Errorf("unexpected query: %s", sql)
}
return q.rows, nil
}

type fakePostgreSQLRows struct {
values []any
read bool
closed bool
err error
}

func (r *fakePostgreSQLRows) Close() {
r.closed = true
}

func (r *fakePostgreSQLRows) Err() error {
return r.err
}

func (r *fakePostgreSQLRows) CommandTag() pgconn.CommandTag {
return pgconn.CommandTag{}
}

func (r *fakePostgreSQLRows) FieldDescriptions() []pgconn.FieldDescription {
return nil
}

func (r *fakePostgreSQLRows) Next() bool {
if r.read {
r.Close()
return false
}
r.read = true
return true
}

func (r *fakePostgreSQLRows) Scan(dest ...any) error {
return fakePostgreSQLRow{values: r.values}.Scan(dest...)
}

func (r *fakePostgreSQLRows) Values() ([]any, error) {
return r.values, nil
}

func (r *fakePostgreSQLRows) RawValues() [][]byte {
return nil
}

func (r *fakePostgreSQLRows) Conn() *pgx.Conn {
return nil
}

func postgreSQLAuditLogRowValues(errorType any) []any {
return []any{
"entry-null-error-type",
time.Unix(1700000000, 0).UTC(),
int64(1234),
"gpt-4o-mini",
"gpt-4o-mini",
"openai",
"primary-openai",
false,
nil,
nil,
200,
"req-1",
nil,
"master_key",
"127.0.0.1",
"POST",
"/v1/chat/completions",
"/",
false,
errorType,
`{"user_agent":"test-agent"}`,
}
}

func TestPostgreSQLReaderGetLogsAllowsNullErrorType(t *testing.T) {
rows := &fakePostgreSQLRows{values: postgreSQLAuditLogRowValues(nil)}
reader := &PostgreSQLReader{
pool: fakePostgreSQLQueryer{
count: 1,
rows: rows,
},
}

result, err := reader.GetLogs(context.Background(), LogQueryParams{Limit: 10})
if err != nil {
t.Fatalf("GetLogs failed: %v", err)
}
if result.Total != 1 {
t.Fatalf("Total = %d, want 1", result.Total)
}
if len(result.Entries) != 1 {
t.Fatalf("len(Entries) = %d, want 1", len(result.Entries))
}
entry := result.Entries[0]
if entry.ErrorType != "" {
t.Fatalf("ErrorType = %q, want empty", entry.ErrorType)
}
if entry.ProviderName != "primary-openai" {
t.Fatalf("ProviderName = %q, want primary-openai", entry.ProviderName)
}
if entry.Data == nil || entry.Data.UserAgent != "test-agent" {
t.Fatalf("Data = %#v, want user_agent", entry.Data)
}
if !rows.closed {
t.Fatal("rows were not closed")
}
}

func TestScanPostgreSQLLogEntryAllowsNullErrorType(t *testing.T) {
entry, err := scanPostgreSQLLogEntry(fakePostgreSQLRow{values: postgreSQLAuditLogRowValues(nil)})
if err != nil {
t.Fatalf("scanPostgreSQLLogEntry failed: %v", err)
}
if entry.ErrorType != "" {
t.Fatalf("ErrorType = %q, want empty", entry.ErrorType)
}
if entry.ProviderName != "primary-openai" {
t.Fatalf("ProviderName = %q, want primary-openai", entry.ProviderName)
}
if entry.AuthMethod != "master_key" {
t.Fatalf("AuthMethod = %q, want master_key", entry.AuthMethod)
}
if entry.Data == nil || entry.Data.UserAgent != "test-agent" {
t.Fatalf("Data = %#v, want user_agent", entry.Data)
}
}
12 changes: 10 additions & 2 deletions internal/auditlog/reader_sqlite.go
Original file line number Diff line number Diff line change
Expand Up @@ -113,9 +113,10 @@ func (r *SQLiteReader) GetLogs(ctx context.Context, params LogQueryParams) (*Log
var authKeyID sql.NullString
var authMethod sql.NullString
var userPath sql.NullString
var errorType sql.NullString

if err := rows.Scan(&e.ID, &ts, &e.DurationNs, &e.RequestedModel, &e.ResolvedModel, &e.Provider, &providerName, &aliasUsedInt, &workflowVersionID, &cacheType, &e.StatusCode,
&e.RequestID, &authKeyID, &authMethod, &e.ClientIP, &e.Method, &e.Path, &userPath, &streamInt, &e.ErrorType, &dataJSON); err != nil {
&e.RequestID, &authKeyID, &authMethod, &e.ClientIP, &e.Method, &e.Path, &userPath, &streamInt, &errorType, &dataJSON); err != nil {
return nil, fmt.Errorf("failed to scan audit log row: %w", err)
}

Expand All @@ -142,6 +143,9 @@ func (r *SQLiteReader) GetLogs(ctx context.Context, params LogQueryParams) (*Log
if userPath.Valid {
e.UserPath = userPath.String
}
if errorType.Valid {
e.ErrorType = errorType.String
}

if dataJSON != nil && *dataJSON != "" {
var data LogData
Expand Down Expand Up @@ -343,9 +347,10 @@ func scanSQLiteLogEntry(rows *sql.Rows) (*LogEntry, error) {
var authKeyID sql.NullString
var authMethod sql.NullString
var userPath sql.NullString
var errorType sql.NullString

if err := rows.Scan(&e.ID, &ts, &e.DurationNs, &e.RequestedModel, &e.ResolvedModel, &e.Provider, &providerName, &aliasUsedInt, &workflowVersionID, &cacheType, &e.StatusCode,
&e.RequestID, &authKeyID, &authMethod, &e.ClientIP, &e.Method, &e.Path, &userPath, &streamInt, &e.ErrorType, &dataJSON); err != nil {
&e.RequestID, &authKeyID, &authMethod, &e.ClientIP, &e.Method, &e.Path, &userPath, &streamInt, &errorType, &dataJSON); err != nil {
return nil, fmt.Errorf("failed to scan audit log row: %w", err)
}

Expand All @@ -372,6 +377,9 @@ func scanSQLiteLogEntry(rows *sql.Rows) (*LogEntry, error) {
if userPath.Valid {
e.UserPath = userPath.String
}
if errorType.Valid {
e.ErrorType = errorType.String
}

if dataJSON != nil && *dataJSON != "" {
var data LogData
Expand Down
10 changes: 8 additions & 2 deletions internal/auditlog/store_sqlite_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -279,7 +279,7 @@ func TestSQLiteStore_WriteBatch_PersistsAliasFields(t *testing.T) {
}
}

func TestSQLiteReader_AllowsNullWorkflowVersionID(t *testing.T) {
func TestSQLiteReader_AllowsNullWorkflowVersionIDAndErrorType(t *testing.T) {
db := createTestDB(t)
defer db.Close()

Expand Down Expand Up @@ -310,7 +310,7 @@ func TestSQLiteReader_AllowsNullWorkflowVersionID(t *testing.T) {
"POST",
"/v1/chat/completions",
0,
"",
nil,
nil,
); err != nil {
t.Fatalf("failed to insert audit log row: %v", err)
Expand All @@ -332,6 +332,9 @@ func TestSQLiteReader_AllowsNullWorkflowVersionID(t *testing.T) {
if entry.WorkflowVersionID != "" {
t.Fatalf("WorkflowVersionID = %q, want empty", entry.WorkflowVersionID)
}
if entry.ErrorType != "" {
t.Fatalf("ErrorType = %q, want empty", entry.ErrorType)
}

logs, err := reader.GetLogs(context.Background(), LogQueryParams{Limit: 10})
if err != nil {
Expand All @@ -343,6 +346,9 @@ func TestSQLiteReader_AllowsNullWorkflowVersionID(t *testing.T) {
if logs.Entries[0].WorkflowVersionID != "" {
t.Fatalf("list WorkflowVersionID = %q, want empty", logs.Entries[0].WorkflowVersionID)
}
if logs.Entries[0].ErrorType != "" {
t.Fatalf("list ErrorType = %q, want empty", logs.Entries[0].ErrorType)
}
}

func TestSQLiteReader_GetLogsFiltersByUserPathSubtree(t *testing.T) {
Expand Down
Loading