290 lines
10 KiB
Go
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)
|