0.10.1: 安全与业务逻辑加固、新品牌与部署加固

三轮审查修复(60+ 项),相对远端 main(b536672)的关键变更:
- 安全: 数据面 SSRF 拨号防护(防 DNS rebinding)/上游凭据剥离/登录防枚举
  与锁定态统一/可信代理(X-Forwarded-For)限流加固/会话版本失效机制/
  撤销即时传播/弱密钥拒绝启动/脱敏字节级重写(保签名契约)
- 业务逻辑: 裸 body 上传 panic/bootstrap 审计管线卡死/定价通配符优先级/
  全局工具可见性/调度器停机补跑/TOTP 挑战令牌消费顺序/熔断探针语义/
  >4MB 响应 token 计量/管理员重置密码作废会话 等
- 前端: 新 logo(语枢 AI 网关主题)/Provider 凭据异常警示/删除入口/
  后端错误消息透传/localStorage 敏感数据收敛
- 部署: CREDENTIAL_MASTER_KEY 持久化与弱值拒绝/Provider DELETE 接口/
  nginx 安全头/worker 内存限制
- 新增迁移 000029(key_hash 索引)/000030(usage_daily 币种维度)
This commit is contained in:
2026-08-13 10:50:51 +08:00
parent b536672000
commit 9501751792
136 changed files with 8024 additions and 1476 deletions
+166
View File
@@ -0,0 +1,166 @@
package trace
import (
"errors"
"net/http"
"strconv"
"strings"
"time"
"aigateway.local/core/internal/identity"
"aigateway.local/core/internal/platform/apiresponse"
)
type AdminHTTPHandler struct {
store *Store
identity *identity.Service
mux *http.ServeMux
}
func NewAdminHTTPHandler(store *Store, identityService *identity.Service) *AdminHTTPHandler {
h := &AdminHTTPHandler{store: store, identity: identityService, mux: http.NewServeMux()}
h.mux.HandleFunc("GET /api/v1/admin/traces", h.list)
h.mux.HandleFunc("GET /api/v1/admin/traces/{id}", h.get)
h.mux.HandleFunc("GET /api/v1/admin/agent-sessions", h.listSessions)
return h
}
func (h *AdminHTTPHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) { h.mux.ServeHTTP(w, r) }
func (h *AdminHTTPHandler) require(w http.ResponseWriter, r *http.Request) bool {
account, err := h.identity.Authenticate(r.Context(), identity.KindAdmin, r.Header.Get("Authorization"))
if err != nil {
status := http.StatusUnauthorized
if !errors.Is(err, identity.ErrInvalidSession) && !errors.Is(err, identity.ErrNotFound) {
status = http.StatusServiceUnavailable
}
apiresponse.Error(w, status, "登录状态无效或身份服务暂不可用")
return false
}
if !identity.HasPermission(account, identity.PermissionTraceRead) {
apiresponse.Error(w, http.StatusForbidden, "缺少 LLM Trace 查看权限")
return false
}
return true
}
func (h *AdminHTTPHandler) list(w http.ResponseWriter, r *http.Request) {
if !h.require(w, r) {
return
}
now := time.Now().UTC()
from, err := queryTime(r, "from", now.Add(-24*time.Hour))
if err != nil {
apiresponse.Error(w, http.StatusBadRequest, "from 时间无效")
return
}
to, err := queryTime(r, "to", now.Add(time.Second))
if err != nil || !to.After(from) || to.Sub(from) > 366*24*time.Hour {
apiresponse.Error(w, http.StatusBadRequest, "Trace 查询时间范围无效或超过 366 天")
return
}
limit := 50
if value := strings.TrimSpace(r.URL.Query().Get("limit")); value != "" {
parsed, scanErr := strconv.Atoi(value)
if scanErr != nil || parsed < 1 || parsed > 200 {
apiresponse.Error(w, http.StatusBadRequest, "limit 必须在 1 到 200 之间")
return
}
limit = parsed
}
status := strings.TrimSpace(r.URL.Query().Get("status"))
if status != "" && !validStatus(status) {
apiresponse.Error(w, http.StatusBadRequest, "Trace 状态无效")
return
}
traceType := strings.TrimSpace(r.URL.Query().Get("trace_type"))
if traceType != "" && !validTraceType(traceType) {
apiresponse.Error(w, http.StatusBadRequest, "Trace 类型无效")
return
}
items, err := h.store.List(r.Context(), Filter{From: from, To: to, TraceType: traceType, TargetCode: strings.TrimSpace(r.URL.Query().Get("target_code")), RequestID: strings.TrimSpace(r.URL.Query().Get("request_id")), Status: status, Limit: limit})
if err != nil {
apiresponse.Error(w, http.StatusServiceUnavailable, "Trace 查询服务暂不可用")
return
}
apiresponse.OK(w, map[string]any{"items": items})
}
func (h *AdminHTTPHandler) listSessions(w http.ResponseWriter, r *http.Request) {
if !h.require(w, r) {
return
}
now := time.Now().UTC()
from, err := queryTime(r, "from", now.Add(-30*24*time.Hour))
if err != nil {
apiresponse.Error(w, http.StatusBadRequest, "from 时间无效")
return
}
to, err := queryTime(r, "to", now.Add(time.Second))
if err != nil || !to.After(from) || to.Sub(from) > 366*24*time.Hour {
apiresponse.Error(w, http.StatusBadRequest, "会话查询时间范围无效或超过 366 天")
return
}
limit := 50
if value := strings.TrimSpace(r.URL.Query().Get("limit")); value != "" {
parsed, scanErr := strconv.Atoi(value)
if scanErr != nil || parsed < 1 || parsed > 200 {
apiresponse.Error(w, http.StatusBadRequest, "limit 必须在 1 到 200 之间")
return
}
limit = parsed
}
traceType := strings.TrimSpace(r.URL.Query().Get("trace_type"))
if traceType != "" && !validTraceType(traceType) {
apiresponse.Error(w, http.StatusBadRequest, "会话类型无效")
return
}
targetCode := strings.TrimSpace(r.URL.Query().Get("target_code"))
if len(targetCode) > 128 {
apiresponse.Error(w, http.StatusBadRequest, "目标编码过长")
return
}
sessionID := strings.TrimSpace(r.URL.Query().Get("session_id"))
if len(sessionID) > 512 {
apiresponse.Error(w, http.StatusBadRequest, "会话 ID 过长")
return
}
items, err := h.store.ListSessions(r.Context(), SessionFilter{From: from, To: to, TraceType: traceType, TargetCode: targetCode, SessionID: sessionID, Limit: limit})
if err != nil {
apiresponse.Error(w, http.StatusServiceUnavailable, "智能体会话查询服务暂不可用")
return
}
apiresponse.OK(w, map[string]any{"items": items})
}
func (h *AdminHTTPHandler) get(w http.ResponseWriter, r *http.Request) {
if !h.require(w, r) {
return
}
item, err := h.store.Get(r.Context(), r.PathValue("id"))
if errors.Is(err, ErrNotFound) {
apiresponse.Error(w, http.StatusNotFound, "Trace 不存在")
return
}
if err != nil {
apiresponse.Error(w, http.StatusServiceUnavailable, "Trace 查询服务暂不可用")
return
}
apiresponse.OK(w, item)
}
func queryTime(r *http.Request, name string, fallback time.Time) (time.Time, error) {
value := strings.TrimSpace(r.URL.Query().Get(name))
if value == "" {
return fallback, nil
}
return time.Parse(time.RFC3339, value)
}
func validTraceType(value string) bool {
return value == "application" || value == "digital_employee"
}
func validStatus(value string) bool {
return value == "running" || value == "success" || value == "error"
}
+459
View File
@@ -0,0 +1,459 @@
package trace
import (
"context"
"encoding/json"
"errors"
"fmt"
"strconv"
"strings"
"time"
"unicode/utf8"
platformid "aigateway.local/core/internal/platform/id"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgxpool"
)
var ErrNotFound = errors.New("trace not found")
type Store struct{ pool *pgxpool.Pool }
func NewStore(pool *pgxpool.Pool) *Store { return &Store{pool: pool} }
type StartInput struct {
RequestID string
APIKeyID string
TenantID *string
TraceType string
TargetID string
TargetCode string
ConversationID string
Metadata map[string]any
}
type Trace struct {
ID string `json:"id"`
RequestID string `json:"request_id"`
APIKeyID *string `json:"api_key_id,omitempty"`
TenantID *string `json:"tenant_id,omitempty"`
TraceType string `json:"trace_type"`
TargetID *string `json:"target_id,omitempty"`
TargetCode string `json:"target_code"`
ConversationID string `json:"conversation_id"`
Status string `json:"status"`
StartedAt time.Time `json:"started_at"`
FinishedAt *time.Time `json:"finished_at,omitempty"`
LatencyMS *int `json:"latency_ms,omitempty"`
RetrievalCount int `json:"retrieval_count"`
ModelCallCount int `json:"model_call_count"`
ToolCallCount int `json:"tool_call_count"`
Error string `json:"error"`
Metadata json.RawMessage `json:"metadata"`
Spans []Span `json:"spans,omitempty"`
}
type SpanInput struct {
TraceID string
ParentID string
SpanType string
Name string
Round int
ProviderCode string
Model string
Metadata map[string]any
}
type Span struct {
ID string `json:"id"`
TraceID string `json:"trace_id"`
ParentID *string `json:"parent_id,omitempty"`
SpanType string `json:"span_type"`
Name string `json:"name"`
Status string `json:"status"`
StartedAt time.Time `json:"started_at"`
FinishedAt *time.Time `json:"finished_at,omitempty"`
LatencyMS *int `json:"latency_ms,omitempty"`
ProviderCode string `json:"provider_code,omitempty"`
Model string `json:"model,omitempty"`
InputTokens int64 `json:"input_tokens"`
OutputTokens int64 `json:"output_tokens"`
Round int `json:"round"`
Error string `json:"error"`
Metadata json.RawMessage `json:"metadata"`
}
type FinishInput struct {
Status string
Error string
RetrievalCount int
ModelCallCount int
ToolCallCount int
Metadata map[string]any
}
type SpanFinishInput struct {
Status string
Error string
InputTokens int64
OutputTokens int64
ProviderCode string
Model string
Metadata map[string]any
}
type Filter struct {
From time.Time
To time.Time
TraceType string
TargetCode string
RequestID string
Status string
Limit int
}
// Session is a metadata-only aggregation of traces that share a conversation
// ID. Stateless requests use a request-derived key so they remain visible in
// the session center without pretending to be part of a persistent chat.
type Session struct {
ID string `json:"id"`
TraceType string `json:"trace_type"`
TargetCode string `json:"target_code"`
TraceCount int `json:"trace_count"`
LatestTraceID string `json:"latest_trace_id"`
LatestStatus string `json:"latest_status"`
StartedAt time.Time `json:"started_at"`
UpdatedAt time.Time `json:"updated_at"`
RetrievalCount int `json:"retrieval_count"`
ModelCallCount int `json:"model_call_count"`
ToolCallCount int `json:"tool_call_count"`
}
type SessionFilter struct {
From time.Time
To time.Time
TraceType string
TargetCode string
SessionID string
Limit int
}
func metadataJSON(value map[string]any) []byte {
if value == nil {
return []byte(`{}`)
}
raw, err := json.Marshal(value)
if err != nil {
return []byte(`{}`)
}
return raw
}
func normalizeMetadata(raw []byte) json.RawMessage {
if len(raw) == 0 || !json.Valid(raw) {
return json.RawMessage(`{}`)
}
return raw
}
func validateStart(input StartInput) error {
if strings.TrimSpace(input.RequestID) == "" || strings.TrimSpace(input.TargetCode) == "" {
return errors.New("trace request_id 和 target_code 不能为空")
}
if input.TraceType != "application" && input.TraceType != "digital_employee" {
return errors.New("trace 类型无效")
}
if len(input.TargetCode) > 128 || len(input.ConversationID) > 128 {
return errors.New("trace 目标或会话 ID 过长")
}
return nil
}
func validateSpan(input SpanInput) error {
if strings.TrimSpace(input.TraceID) == "" || strings.TrimSpace(input.Name) == "" {
return errors.New("trace span 标识不能为空")
}
if input.SpanType != "model" && input.SpanType != "tool" && input.SpanType != "retrieval" {
return errors.New("trace span 类型无效")
}
if input.Round < 0 || len(input.Name) > 256 {
return errors.New("trace span 参数无效")
}
return nil
}
func (s *Store) Start(ctx context.Context, input StartInput) (Trace, error) {
if s == nil || s.pool == nil {
return Trace{}, errors.New("trace store unavailable")
}
input.RequestID = strings.TrimSpace(input.RequestID)
input.TargetCode = strings.TrimSpace(input.TargetCode)
input.ConversationID = strings.TrimSpace(input.ConversationID)
if err := validateStart(input); err != nil {
return Trace{}, err
}
id, err := platformid.NewUUID()
if err != nil {
return Trace{}, err
}
_, err = s.pool.Exec(ctx, `INSERT INTO gateway.agent_traces(id,request_id,api_key_id,tenant_id,trace_type,target_id,target_code,conversation_id,metadata) VALUES($1,$2,nullif($3,'')::uuid,nullif($4,'')::uuid,$5,nullif($6,'')::uuid,$7,$8,$9)`, id, input.RequestID, input.APIKeyID, valueOrEmpty(input.TenantID), input.TraceType, input.TargetID, input.TargetCode, input.ConversationID, metadataJSON(input.Metadata))
if err != nil {
return Trace{}, fmt.Errorf("start trace: %w", err)
}
return s.Get(ctx, id)
}
func valueOrEmpty(value *string) string {
if value == nil {
return ""
}
return *value
}
func (s *Store) StartSpan(ctx context.Context, input SpanInput) (Span, error) {
if s == nil || s.pool == nil {
return Span{}, errors.New("trace store unavailable")
}
input.Name = strings.TrimSpace(input.Name)
if err := validateSpan(input); err != nil {
return Span{}, err
}
id, err := platformid.NewUUID()
if err != nil {
return Span{}, err
}
_, err = s.pool.Exec(ctx, `INSERT INTO gateway.agent_trace_spans(id,trace_id,parent_id,span_type,name,round,provider_code,model,metadata) VALUES($1,$2,nullif($3,'')::uuid,$4,$5,$6,nullif($7,''),nullif($8,''),$9)`, id, input.TraceID, input.ParentID, input.SpanType, input.Name, input.Round, input.ProviderCode, input.Model, metadataJSON(input.Metadata))
if err != nil {
return Span{}, fmt.Errorf("start trace span: %w", err)
}
return s.GetSpan(ctx, id)
}
func (s *Store) Finish(ctx context.Context, id string, input FinishInput) error {
if s == nil || s.pool == nil {
return errors.New("trace store unavailable")
}
status := normalizeStatus(input.Status)
errorText := truncate(input.Error, 4000)
metadata := metadataJSON(input.Metadata)
tag, err := s.pool.Exec(ctx, `UPDATE gateway.agent_traces SET status=$2,error=$3,retrieval_count=$4,model_call_count=$5,tool_call_count=$6,metadata=$7,finished_at=clock_timestamp(),latency_ms=(extract(epoch FROM (clock_timestamp()-started_at))*1000)::integer WHERE id=$1 AND status='running'`, id, status, errorText, max(input.RetrievalCount, 0), max(input.ModelCallCount, 0), max(input.ToolCallCount, 0), metadata)
if err != nil {
return err
}
if tag.RowsAffected() == 0 {
return ErrNotFound
}
return nil
}
func (s *Store) FinishSpan(ctx context.Context, id string, input SpanFinishInput) error {
if s == nil || s.pool == nil {
return errors.New("trace store unavailable")
}
status := normalizeStatus(input.Status)
tag, err := s.pool.Exec(ctx, `UPDATE gateway.agent_trace_spans SET status=$2,error=$3,input_tokens=$4,output_tokens=$5,provider_code=coalesce(nullif($6,''),provider_code),model=coalesce(nullif($7,''),model),metadata=$8,finished_at=clock_timestamp(),latency_ms=(extract(epoch FROM (clock_timestamp()-started_at))*1000)::integer WHERE id=$1 AND status='running'`, id, status, truncate(input.Error, 4000), max(input.InputTokens, 0), max(input.OutputTokens, 0), input.ProviderCode, input.Model, metadataJSON(input.Metadata))
if err != nil {
return err
}
if tag.RowsAffected() == 0 {
return ErrNotFound
}
return nil
}
func normalizeStatus(value string) string {
if value == "success" {
return "success"
}
return "error"
}
func truncate(value string, limit int) string {
value = strings.TrimSpace(value)
if len(value) <= limit {
return value
}
cut := value[:limit]
// 按字节截断可能切半多字节 rune,产生无效 UTF-8 使 trace 落库失败;
// 回退到最近一个完整 rune 的边界。
for len(cut) > 0 && !utf8.RuneStart(cut[len(cut)-1]) {
cut = cut[:len(cut)-1]
}
return cut
}
const traceSelect = `SELECT id::text,request_id,api_key_id::text,tenant_id::text,trace_type,target_id::text,target_code,conversation_id,status,started_at,finished_at,latency_ms,retrieval_count,model_call_count,tool_call_count,error,metadata FROM gateway.agent_traces`
func scanTrace(row pgx.Row) (Trace, error) {
var item Trace
err := row.Scan(&item.ID, &item.RequestID, &item.APIKeyID, &item.TenantID, &item.TraceType, &item.TargetID, &item.TargetCode, &item.ConversationID, &item.Status, &item.StartedAt, &item.FinishedAt, &item.LatencyMS, &item.RetrievalCount, &item.ModelCallCount, &item.ToolCallCount, &item.Error, &item.Metadata)
if errors.Is(err, pgx.ErrNoRows) {
return Trace{}, ErrNotFound
}
item.Metadata = normalizeMetadata(item.Metadata)
return item, err
}
func scanSpan(row pgx.Row) (Span, error) {
var item Span
err := row.Scan(&item.ID, &item.TraceID, &item.ParentID, &item.SpanType, &item.Name, &item.Status, &item.StartedAt, &item.FinishedAt, &item.LatencyMS, &item.ProviderCode, &item.Model, &item.InputTokens, &item.OutputTokens, &item.Round, &item.Error, &item.Metadata)
if errors.Is(err, pgx.ErrNoRows) {
return Span{}, ErrNotFound
}
item.Metadata = normalizeMetadata(item.Metadata)
return item, err
}
func (s *Store) Get(ctx context.Context, id string) (Trace, error) {
if s == nil || s.pool == nil {
return Trace{}, errors.New("trace store unavailable")
}
item, err := scanTrace(s.pool.QueryRow(ctx, traceSelect+` WHERE id=$1`, id))
if err != nil {
return Trace{}, err
}
spans, err := s.listSpans(ctx, id)
if err != nil {
return Trace{}, err
}
item.Spans = spans
return item, nil
}
func (s *Store) GetSpan(ctx context.Context, id string) (Span, error) {
if s == nil || s.pool == nil {
return Span{}, errors.New("trace store unavailable")
}
return scanSpan(s.pool.QueryRow(ctx, `SELECT id::text,trace_id::text,parent_id::text,span_type,name,status,started_at,finished_at,latency_ms,coalesce(provider_code,''),coalesce(model,''),input_tokens,output_tokens,round,error,metadata FROM gateway.agent_trace_spans WHERE id=$1`, id))
}
func (s *Store) listSpans(ctx context.Context, traceID string) ([]Span, error) {
if s == nil || s.pool == nil {
return nil, errors.New("trace store unavailable")
}
rows, err := s.pool.Query(ctx, `SELECT id::text,trace_id::text,parent_id::text,span_type,name,status,started_at,finished_at,latency_ms,coalesce(provider_code,''),coalesce(model,''),input_tokens,output_tokens,round,error,metadata FROM gateway.agent_trace_spans WHERE trace_id=$1 ORDER BY started_at,id`, traceID)
if err != nil {
return nil, err
}
defer rows.Close()
items := []Span{}
for rows.Next() {
item, scanErr := scanSpan(rows)
if scanErr != nil {
return nil, scanErr
}
items = append(items, item)
}
return items, rows.Err()
}
func (s *Store) List(ctx context.Context, filter Filter) ([]Trace, error) {
if s == nil || s.pool == nil {
return nil, errors.New("trace store unavailable")
}
if filter.Limit < 1 || filter.Limit > 200 {
filter.Limit = 50
}
where := []string{"started_at >= $1", "started_at < $2"}
args := []any{filter.From, filter.To}
add := func(condition string, value any) {
args = append(args, value)
where = append(where, fmt.Sprintf(condition, len(args)))
}
if filter.TraceType != "" {
add("trace_type = $%d", filter.TraceType)
}
if filter.TargetCode != "" {
add("target_code = $%d", filter.TargetCode)
}
if filter.RequestID != "" {
add("request_id = $%d", filter.RequestID)
}
if filter.Status != "" {
add("status = $%d", filter.Status)
}
args = append(args, filter.Limit)
query := traceSelect + ` WHERE ` + strings.Join(where, " AND ") + ` ORDER BY started_at DESC,id DESC LIMIT $` + strconv.Itoa(len(args))
rows, err := s.pool.Query(ctx, query, args...)
if err != nil {
return nil, err
}
defer rows.Close()
items := []Trace{}
for rows.Next() {
item, scanErr := scanTrace(rows)
if scanErr != nil {
return nil, scanErr
}
items = append(items, item)
}
return items, rows.Err()
}
const sessionGroupSelect = `
WITH grouped AS (
SELECT trace_type,
target_code,
coalesce(nullif(conversation_id, ''), 'request:' || request_id) AS session_key,
count(*)::int AS trace_count,
(array_agg(id::text ORDER BY started_at DESC, id DESC))[1] AS latest_trace_id,
(array_agg(status ORDER BY started_at DESC, id DESC))[1] AS latest_status,
min(started_at) AS started_at,
max(coalesce(finished_at, started_at)) AS updated_at,
sum(retrieval_count)::int AS retrieval_count,
sum(model_call_count)::int AS model_call_count,
sum(tool_call_count)::int AS tool_call_count
FROM gateway.agent_traces`
func scanSession(row pgx.Row) (Session, error) {
var item Session
err := row.Scan(&item.ID, &item.TraceType, &item.TargetCode, &item.TraceCount, &item.LatestTraceID, &item.LatestStatus, &item.StartedAt, &item.UpdatedAt, &item.RetrievalCount, &item.ModelCallCount, &item.ToolCallCount)
return item, err
}
func (s *Store) ListSessions(ctx context.Context, filter SessionFilter) ([]Session, error) {
if s == nil || s.pool == nil {
return nil, errors.New("trace store unavailable")
}
if filter.Limit < 1 || filter.Limit > 200 {
filter.Limit = 50
}
innerWhere := []string{"started_at >= $1", "started_at < $2"}
args := []any{filter.From, filter.To}
addInner := func(condition string, value any) {
args = append(args, value)
innerWhere = append(innerWhere, fmt.Sprintf(condition, len(args)))
}
if filter.TraceType != "" {
addInner("trace_type = $%d", filter.TraceType)
}
if filter.TargetCode != "" {
addInner("target_code = $%d", filter.TargetCode)
}
outerWhere := []string{}
if filter.SessionID != "" {
args = append(args, filter.SessionID)
outerWhere = append(outerWhere, fmt.Sprintf("trace_type || ':' || target_code || ':' || session_key = $%d", len(args)))
}
args = append(args, filter.Limit)
limitArg := strconv.Itoa(len(args))
query := sessionGroupSelect + ` WHERE ` + strings.Join(innerWhere, " AND ") + ` GROUP BY trace_type,target_code,session_key) SELECT trace_type || ':' || target_code || ':' || session_key AS id,trace_type,target_code,trace_count,latest_trace_id,latest_status,started_at,updated_at,retrieval_count,model_call_count,tool_call_count FROM grouped`
if len(outerWhere) > 0 {
query += ` WHERE ` + strings.Join(outerWhere, " AND ")
}
query += ` ORDER BY updated_at DESC,id DESC LIMIT $` + limitArg
rows, err := s.pool.Query(ctx, query, args...)
if err != nil {
return nil, err
}
defer rows.Close()
items := []Session{}
for rows.Next() {
item, scanErr := scanSession(rows)
if scanErr != nil {
return nil, scanErr
}
items = append(items, item)
}
return items, rows.Err()
}
+87
View File
@@ -0,0 +1,87 @@
package trace
import (
"context"
"os"
"strings"
"testing"
"time"
"aigateway.local/core/internal/platform/config"
"aigateway.local/core/internal/platform/database"
)
func TestTracePostgreSQLLifecycle(t *testing.T) {
databaseURL := os.Getenv("TRACE_TEST_DATABASE_URL")
if databaseURL == "" {
t.Skip("TRACE_TEST_DATABASE_URL is not set")
}
ctx := context.Background()
pool, err := database.Open(ctx, config.Database{URL: databaseURL, MaxConns: 4, MinConns: 0})
if err != nil {
t.Fatal(err)
}
defer pool.Close()
store := NewStore(pool)
_, _ = pool.Exec(ctx, `DELETE FROM gateway.agent_traces WHERE request_id LIKE 'trace-integration%'`)
item, err := store.Start(ctx, StartInput{RequestID: "trace-integration", TraceType: "application", TargetID: "77777777-7777-4777-8777-777777777777", TargetCode: "trace_app", ConversationID: "conversation-1", Metadata: map[string]any{"version": 1}})
if err != nil {
t.Fatal(err)
}
defer pool.Exec(ctx, `DELETE FROM gateway.agent_traces WHERE request_id LIKE 'trace-integration%'`)
model, err := store.StartSpan(ctx, SpanInput{TraceID: item.ID, SpanType: "model", Name: "chat.completions", Round: 0, ProviderCode: "test", Model: "test-model"})
if err != nil {
t.Fatal(err)
}
if err = store.FinishSpan(ctx, model.ID, SpanFinishInput{Status: "success", InputTokens: 12, OutputTokens: 8, Metadata: map[string]any{"http_status": 200}}); err != nil {
t.Fatal(err)
}
tool, err := store.StartSpan(ctx, SpanInput{TraceID: item.ID, SpanType: "tool", Name: "lookup", Round: 0})
if err != nil {
t.Fatal(err)
}
if err = store.FinishSpan(ctx, tool.ID, SpanFinishInput{Status: "error", Error: "upstream timeout"}); err != nil {
t.Fatal(err)
}
if err = store.Finish(ctx, item.ID, FinishInput{Status: "success", RetrievalCount: 2, ModelCallCount: 1, ToolCallCount: 1}); err != nil {
t.Fatal(err)
}
loaded, err := store.Get(ctx, item.ID)
if err != nil || loaded.Status != "success" || loaded.ModelCallCount != 1 || len(loaded.Spans) != 2 {
t.Fatalf("loaded=%+v err=%v", loaded, err)
}
if loaded.Spans[0].ProviderCode != "test" || loaded.Spans[0].Model != "test-model" {
t.Fatalf("model span route metadata was not preserved: %+v", loaded.Spans[0])
}
if loaded.LatencyMS == nil || loaded.StartedAt.After(time.Now().UTC().Add(time.Second)) {
t.Fatalf("invalid timing: %+v", loaded)
}
items, err := store.List(ctx, Filter{From: time.Now().Add(-time.Minute), To: time.Now().Add(time.Minute), TargetCode: "trace_app", Limit: 10})
if err != nil || len(items) != 1 || items[0].ID != item.ID {
t.Fatalf("list=%+v err=%v", items, err)
}
sessions, err := store.ListSessions(ctx, SessionFilter{From: time.Now().Add(-time.Minute), To: time.Now().Add(time.Minute), SessionID: "application:trace_app:conversation-1", Limit: 10})
if err != nil || len(sessions) != 1 || sessions[0].TraceCount != 1 || sessions[0].LatestTraceID != item.ID || sessions[0].ModelCallCount != 1 {
t.Fatalf("sessions=%+v err=%v", sessions, err)
}
digital, err := store.Start(ctx, StartInput{RequestID: "trace-integration-digital", TraceType: "digital_employee", TargetID: "88888888-8888-4888-8888-888888888888", TargetCode: "trace_employee", ConversationID: "conversation-1"})
if err != nil {
t.Fatal(err)
}
if err = store.Finish(ctx, digital.ID, FinishInput{Status: "success", ModelCallCount: 1}); err != nil {
t.Fatal(err)
}
allSessions, err := store.ListSessions(ctx, SessionFilter{From: time.Now().Add(-time.Minute), To: time.Now().Add(time.Minute), Limit: 10})
if err != nil || len(allSessions) < 2 {
t.Fatalf("all sessions=%+v err=%v", allSessions, err)
}
seenTypes := map[string]bool{}
for _, session := range allSessions {
if strings.HasPrefix(session.ID, "application:trace_app:") || strings.HasPrefix(session.ID, "digital_employee:trace_employee:") {
seenTypes[session.TraceType] = true
}
}
if !seenTypes["application"] || !seenTypes["digital_employee"] {
t.Fatalf("session types=%v", seenTypes)
}
}