Files
chorus/portal/worker/runtime.go
T

290 lines
10 KiB
Go

package worker
import (
"context"
"crypto/rand"
"encoding/hex"
"fmt"
"strings"
"time"
"git.ilapage.cn/OPC/chorus/internal/core/model"
"git.ilapage.cn/OPC/chorus/internal/core/queue"
"git.ilapage.cn/OPC/chorus/internal/core/router"
"gorm.io/gorm"
)
// GORMRuntime is the cross-worker circuit breaker. It intentionally queries
// only runtime and enabled/capability flags; credentials and provider URLs are
// never part of this read path.
type GORMRuntime struct{ db *gorm.DB }
func NewGORMRuntime(db *gorm.DB) (*GORMRuntime, error) {
if db == nil {
return nil, ErrInvalidConfig
}
return &GORMRuntime{db: db}, nil
}
func (r *GORMRuntime) MemberStates(ctx context.Context, capability model.Capability, members []router.MemberSnapshot) ([]router.MemberState, error) {
if !capability.Valid() || len(members) == 0 {
return nil, router.ErrInvalidSnapshot
}
memberIDs := make([]uint64, 0, len(members))
states := make(map[uint64]router.MemberState, len(members))
for _, member := range members {
if member.RoutePoolMemberID == 0 {
return nil, router.ErrInvalidSnapshot
}
memberIDs = append(memberIDs, member.RoutePoolMemberID)
states[member.RoutePoolMemberID] = router.MemberState{RoutePoolMemberID: member.RoutePoolMemberID}
}
type row struct {
RoutePoolMemberID uint64
MemberEnabled bool
ProviderEnabled bool
ModelEnabled bool
SupportsCapability bool
CircuitState string
OpenUntil *time.Time
ProbeUntil *time.Time
HalfOpenMax uint16
}
var rows []row
err := r.db.WithContext(ctx).Raw(`
SELECT rpm.id AS route_pool_member_id,
rpm.enabled AS member_enabled,
p.enabled AS provider_enabled,
pm.enabled AS model_enabled,
EXISTS(SELECT 1 FROM provider_model_capabilities c
WHERE c.provider_model_id = pm.id AND c.capability = ?) AS supports_capability,
COALESCE(runtime.state, 'closed') AS circuit_state,
runtime.open_until, runtime.probe_until, rpm.half_open_max
FROM route_pool_members rpm
JOIN provider_models pm ON pm.id = rpm.provider_model_id
JOIN providers p ON p.id = pm.provider_id
LEFT JOIN route_member_runtime runtime ON runtime.route_pool_member_id = rpm.id
WHERE rpm.id IN ?`, capability, memberIDs).Scan(&rows).Error
if err != nil {
return nil, fmt.Errorf("load route member states: %w", err)
}
now := time.Now().UTC()
for _, row := range rows {
state := router.CircuitState(row.CircuitState)
halfOpenInFlight := uint16(0)
switch state {
case router.CircuitClosed:
case router.CircuitOpen:
if row.OpenUntil != nil && !row.OpenUntil.After(now) {
state = router.CircuitHalfOpen
}
case router.CircuitHalfOpen:
if row.ProbeUntil != nil && row.ProbeUntil.After(now) {
halfOpenInFlight = 1
}
default:
state = router.CircuitOpen
}
states[row.RoutePoolMemberID] = router.MemberState{
RoutePoolMemberID: row.RoutePoolMemberID, Enabled: row.MemberEnabled,
ProviderEnabled: row.ProviderEnabled, ModelEnabled: row.ModelEnabled,
SupportsCapability: row.SupportsCapability, CircuitState: state,
HalfOpenInFlight: halfOpenInFlight, HalfOpenMax: row.HalfOpenMax,
}
}
result := make([]router.MemberState, 0, len(memberIDs))
for _, memberID := range memberIDs {
result = append(result, states[memberID])
}
return result, nil
}
func (r *GORMRuntime) Reserve(ctx context.Context, request router.ReservationRequest) (router.Reservation, error) {
if request.RoutePoolMemberID == 0 || !request.Capability.Valid() || strings.TrimSpace(request.Owner) == "" || len(request.Owner) > 128 || request.LeaseDuration <= 0 {
return router.Reservation{}, router.ErrMemberUnavailable
}
var reservation router.Reservation
err := r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Exec(`INSERT IGNORE INTO route_member_runtime (route_pool_member_id, state, consecutive_failures) VALUES (?, 'closed', 0)`, request.RoutePoolMemberID).Error; err != nil {
return fmt.Errorf("create route member runtime: %w", err)
}
type row struct {
MemberEnabled bool
ProviderEnabled bool
ModelEnabled bool
SupportsCapability bool
State string
OpenUntil *time.Time
ProbeToken *string
ProbeUntil *time.Time
}
var value row
result := tx.Raw(`
SELECT rpm.enabled AS member_enabled, p.enabled AS provider_enabled, pm.enabled AS model_enabled,
EXISTS(SELECT 1 FROM provider_model_capabilities c
WHERE c.provider_model_id = pm.id AND c.capability = ?) AS supports_capability,
runtime.state, runtime.open_until, runtime.probe_token, runtime.probe_until
FROM route_pool_members rpm
JOIN provider_models pm ON pm.id = rpm.provider_model_id
JOIN providers p ON p.id = pm.provider_id
JOIN route_member_runtime runtime ON runtime.route_pool_member_id = rpm.id
WHERE rpm.id = ? FOR UPDATE`, request.Capability, request.RoutePoolMemberID).Scan(&value)
if result.Error != nil {
return fmt.Errorf("lock route member runtime: %w", result.Error)
}
if result.RowsAffected != 1 || !value.MemberEnabled || !value.ProviderEnabled || !value.ModelEnabled || !value.SupportsCapability {
return router.ErrMemberUnavailable
}
now := time.Now().UTC()
state := router.CircuitState(value.State)
switch state {
case router.CircuitClosed:
reservation = router.Reservation{RoutePoolMemberID: request.RoutePoolMemberID, State: router.CircuitClosed}
return nil
case router.CircuitOpen:
if value.OpenUntil == nil || value.OpenUntil.After(now) {
return router.ErrMemberUnavailable
}
case router.CircuitHalfOpen:
if value.ProbeUntil != nil && value.ProbeUntil.After(now) {
return router.ErrMemberUnavailable
}
default:
return router.ErrMemberUnavailable
}
token, tokenErr := newRuntimeToken()
if tokenErr != nil {
return tokenErr
}
until := now.Add(request.LeaseDuration)
update := tx.Model(&runtimeRow{}).Where("route_pool_member_id = ?", request.RoutePoolMemberID).Updates(map[string]any{
"state": router.CircuitHalfOpen, "open_until": nil, "probe_owner": request.Owner,
"probe_token": token, "probe_until": until,
})
if update.Error != nil {
return fmt.Errorf("reserve half-open route member: %w", update.Error)
}
if update.RowsAffected != 1 {
return router.ErrMemberUnavailable
}
reservation = router.Reservation{RoutePoolMemberID: request.RoutePoolMemberID, Token: token, State: router.CircuitHalfOpen}
return nil
})
return reservation, err
}
func (r *GORMRuntime) Record(ctx context.Context, reservation router.Reservation, observation router.CircuitObservation) (bool, error) {
if reservation.RoutePoolMemberID == 0 {
return false, router.ErrMemberUnavailable
}
owned := false
err := r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
type row struct {
State string
ConsecutiveFailures uint32
ProbeToken *string
ProbeUntil *time.Time
FailureThreshold uint32
OpenSeconds uint32
}
var value row
result := tx.Raw(`
SELECT runtime.state, runtime.consecutive_failures, runtime.probe_token, runtime.probe_until,
rpm.failure_threshold, rpm.open_seconds
FROM route_member_runtime runtime
JOIN route_pool_members rpm ON rpm.id = runtime.route_pool_member_id
WHERE runtime.route_pool_member_id = ? FOR UPDATE`, reservation.RoutePoolMemberID).Scan(&value)
if result.Error != nil {
return fmt.Errorf("lock route observation runtime: %w", result.Error)
}
if result.RowsAffected != 1 {
return nil
}
if reservation.State == router.CircuitHalfOpen {
if value.State != string(router.CircuitHalfOpen) || value.ProbeToken == nil || *value.ProbeToken != reservation.Token || value.ProbeUntil == nil || !value.ProbeUntil.After(time.Now().UTC()) {
return nil
}
}
updates := map[string]any{
"last_observed_at": time.Now().UTC(),
"last_error_code": nullableRuntimeCode(observation.ErrorCode),
"last_error_message": nullableRuntimeMessage(observation.ErrorMessage),
}
switch {
case observation.Succeeded:
updates["state"] = router.CircuitClosed
updates["consecutive_failures"] = 0
updates["open_until"] = nil
case observation.OpenImmediately:
updates["state"] = router.CircuitOpen
updates["consecutive_failures"] = value.FailureThreshold
updates["open_until"] = time.Now().UTC().Add(time.Duration(value.OpenSeconds) * time.Second)
case observation.Retryable:
next := value.ConsecutiveFailures + 1
updates["consecutive_failures"] = next
if next >= value.FailureThreshold {
updates["state"] = router.CircuitOpen
updates["open_until"] = time.Now().UTC().Add(time.Duration(value.OpenSeconds) * time.Second)
} else {
updates["state"] = router.CircuitClosed
updates["open_until"] = nil
}
default:
updates["state"] = router.CircuitClosed
updates["consecutive_failures"] = 0
updates["open_until"] = nil
}
updates["probe_owner"] = nil
updates["probe_token"] = nil
updates["probe_until"] = nil
update := tx.Model(&runtimeRow{}).Where("route_pool_member_id = ?", reservation.RoutePoolMemberID).Updates(updates)
if update.Error != nil {
return fmt.Errorf("record route observation: %w", update.Error)
}
owned = update.RowsAffected == 1
return nil
})
return owned, err
}
type runtimeRow struct {
RoutePoolMemberID uint64 `gorm:"column:route_pool_member_id;primaryKey"`
}
func (runtimeRow) TableName() string { return "route_member_runtime" }
func nullableRuntimeCode(value string) any {
value = strings.ToLower(strings.TrimSpace(value))
if value == "" {
return nil
}
if len(value) > 64 {
return "upstream_error"
}
for _, character := range value {
if character != '_' && (character < 'a' || character > 'z') && (character < '0' || character > '9') {
return "upstream_error"
}
}
return value
}
func nullableRuntimeMessage(value string) any {
value = queue.SanitizeErrorMessage(value, 1024)
if value == "upstream request failed" {
return nil
}
return value
}
func newRuntimeToken() (string, error) {
value := make([]byte, 16)
if _, err := rand.Read(value); err != nil {
return "", fmt.Errorf("generate route runtime token: %w", err)
}
hexValue := hex.EncodeToString(value)
return hexValue[0:8] + "-" + hexValue[8:12] + "-" + hexValue[12:16] + "-" + hexValue[16:20] + "-" + hexValue[20:32], nil
}
var _ router.RuntimeRepository = (*GORMRuntime)(nil)