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
18 changes: 18 additions & 0 deletions common/utils.go
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@ import (
"math/big"
"math/rand"
"net"
"net/url"
"os"
"os/exec"
"runtime"
Expand Down Expand Up @@ -284,3 +285,20 @@ func GetAudioDuration(ctx context.Context, filename string, ext string) (float64
}
return strconv.ParseFloat(durationStr, 64)
}

// BuildURL concatenates base and endpoint, returns the complete url string
func BuildURL(base string, endpoint string) string {
u, err := url.Parse(base)
if err != nil {
return base + endpoint
}
end := endpoint
if end == "" {
end = "/"
}
ref, err := url.Parse(end)
if err != nil {
return base + endpoint
}
return u.ResolveReference(ref).String()
}
24 changes: 24 additions & 0 deletions controller/ratio_config.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,24 @@
package controller

import (
"net/http"
"one-api/setting/ratio_setting"

"github.com/gin-gonic/gin"
)

func GetRatioConfig(c *gin.Context) {
if !ratio_setting.IsExposeRatioEnabled() {
c.JSON(http.StatusForbidden, gin.H{
"success": false,
"message": "倍率配置接口未启用",
})
return
}

c.JSON(http.StatusOK, gin.H{
"success": true,
"message": "",
"data": ratio_setting.GetExposedData(),
})
}
322 changes: 322 additions & 0 deletions controller/ratio_sync.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,322 @@
package controller

import (
"context"
"encoding/json"
"net/http"
"strings"
"sync"
"time"

"one-api/common"
"one-api/dto"
"one-api/model"
"one-api/setting/ratio_setting"

"github.com/gin-gonic/gin"
)

const (
defaultTimeoutSeconds = 10
defaultEndpoint = "/api/ratio_config"
maxConcurrentFetches = 8
)

var ratioTypes = []string{"model_ratio", "completion_ratio", "cache_ratio", "model_price"}

type upstreamResult struct {
Name string `json:"name"`
Data map[string]any `json:"data,omitempty"`
Err string `json:"err,omitempty"`
}

func FetchUpstreamRatios(c *gin.Context) {
var req dto.UpstreamRequest
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"success": false, "message": err.Error()})
return
}

if req.Timeout <= 0 {
req.Timeout = defaultTimeoutSeconds
}

var upstreams []dto.UpstreamDTO

if len(req.ChannelIDs) > 0 {
intIds := make([]int, 0, len(req.ChannelIDs))
for _, id64 := range req.ChannelIDs {
intIds = append(intIds, int(id64))
}
dbChannels, err := model.GetChannelsByIds(intIds)
if err != nil {
common.LogError(c.Request.Context(), "failed to query channels: "+err.Error())
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "message": "查询渠道失败"})
return
}
for _, ch := range dbChannels {
if base := ch.GetBaseURL(); strings.HasPrefix(base, "http") {
upstreams = append(upstreams, dto.UpstreamDTO{
Name: ch.Name,
BaseURL: strings.TrimRight(base, "/"),
Endpoint: "",
})
}
}
}

if len(upstreams) == 0 {
c.JSON(http.StatusOK, gin.H{"success": false, "message": "无有效上游渠道"})
return
}

var wg sync.WaitGroup
ch := make(chan upstreamResult, len(upstreams))

sem := make(chan struct{}, maxConcurrentFetches)

client := &http.Client{Transport: &http.Transport{MaxIdleConns: 100, IdleConnTimeout: 90 * time.Second, TLSHandshakeTimeout: 10 * time.Second, ExpectContinueTimeout: 1 * time.Second}}

for _, chn := range upstreams {
wg.Add(1)
go func(chItem dto.UpstreamDTO) {
defer wg.Done()

sem <- struct{}{}
defer func() { <-sem }()

endpoint := chItem.Endpoint
if endpoint == "" {
endpoint = defaultEndpoint
} else if !strings.HasPrefix(endpoint, "/") {
endpoint = "/" + endpoint
}
fullURL := chItem.BaseURL + endpoint

ctx, cancel := context.WithTimeout(c.Request.Context(), time.Duration(req.Timeout)*time.Second)
defer cancel()

httpReq, err := http.NewRequestWithContext(ctx, http.MethodGet, fullURL, nil)
if err != nil {
common.LogWarn(c.Request.Context(), "build request failed: "+err.Error())
ch <- upstreamResult{Name: chItem.Name, Err: err.Error()}
return
}

resp, err := client.Do(httpReq)
if err != nil {
common.LogWarn(c.Request.Context(), "http error on "+chItem.Name+": "+err.Error())
ch <- upstreamResult{Name: chItem.Name, Err: err.Error()}
return
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
common.LogWarn(c.Request.Context(), "non-200 from "+chItem.Name+": "+resp.Status)
ch <- upstreamResult{Name: chItem.Name, Err: resp.Status}
return
}
var body struct {
Success bool `json:"success"`
Data map[string]any `json:"data"`
Message string `json:"message"`
}
if err := json.NewDecoder(resp.Body).Decode(&body); err != nil {
common.LogWarn(c.Request.Context(), "json decode failed from "+chItem.Name+": "+err.Error())
ch <- upstreamResult{Name: chItem.Name, Err: err.Error()}
return
}
if !body.Success {
ch <- upstreamResult{Name: chItem.Name, Err: body.Message}
return
}
ch <- upstreamResult{Name: chItem.Name, Data: body.Data}
}(chn)
}

wg.Wait()
close(ch)

localData := ratio_setting.GetExposedData()

var testResults []dto.TestResult
var successfulChannels []struct {
name string
data map[string]any
}

for r := range ch {
if r.Err != "" {
testResults = append(testResults, dto.TestResult{
Name: r.Name,
Status: "error",
Error: r.Err,
})
} else {
testResults = append(testResults, dto.TestResult{
Name: r.Name,
Status: "success",
})
successfulChannels = append(successfulChannels, struct {
name string
data map[string]any
}{name: r.Name, data: r.Data})
}
}

differences := buildDifferences(localData, successfulChannels)

c.JSON(http.StatusOK, gin.H{
"success": true,
"data": gin.H{
"differences": differences,
"test_results": testResults,
},
})
}

func buildDifferences(localData map[string]any, successfulChannels []struct {
name string
data map[string]any
}) map[string]map[string]dto.DifferenceItem {
differences := make(map[string]map[string]dto.DifferenceItem)

allModels := make(map[string]struct{})

for _, ratioType := range ratioTypes {
if localRatioAny, ok := localData[ratioType]; ok {
if localRatio, ok := localRatioAny.(map[string]float64); ok {
for modelName := range localRatio {
allModels[modelName] = struct{}{}
}
}
}
}

for _, channel := range successfulChannels {
for _, ratioType := range ratioTypes {
if upstreamRatio, ok := channel.data[ratioType].(map[string]any); ok {
for modelName := range upstreamRatio {
allModels[modelName] = struct{}{}
}
}
}
}

for modelName := range allModels {
for _, ratioType := range ratioTypes {
var localValue interface{} = nil
if localRatioAny, ok := localData[ratioType]; ok {
if localRatio, ok := localRatioAny.(map[string]float64); ok {
if val, exists := localRatio[modelName]; exists {
localValue = val
}
}
}

upstreamValues := make(map[string]interface{})
hasUpstreamValue := false
hasDifference := false

for _, channel := range successfulChannels {
var upstreamValue interface{} = nil

if upstreamRatio, ok := channel.data[ratioType].(map[string]any); ok {
if val, exists := upstreamRatio[modelName]; exists {
upstreamValue = val
hasUpstreamValue = true

if localValue != nil && localValue != val {
hasDifference = true
} else if localValue == val {
upstreamValue = "same"
}
Comment on lines +228 to +232

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🛠️ Refactor suggestion

Fix potential type comparison issue.

The comparison localValue != val may not work correctly for floating-point numbers due to precision issues, and the type assertion could fail silently.

+import "math"

+func compareFloatValues(a, b interface{}) bool {
+    aFloat, aOk := a.(float64)
+    bFloat, bOk := b.(float64)
+    if aOk && bOk {
+        return math.Abs(aFloat-bFloat) < 1e-9  // Use epsilon for float comparison
+    }
+    return a == b  // Fallback to direct comparison
+}

-if localValue != nil && localValue != val {
+if localValue != nil && !compareFloatValues(localValue, val) {
    hasDifference = true
-} else if localValue == val {
+} else if compareFloatValues(localValue, val) {
    upstreamValue = "same"
}
🤖 Prompt for AI Agents
In controller/ratio_sync.go around lines 199 to 203, the code compares
localValue and val directly, which can cause issues with floating-point
precision and silent type assertion failures. To fix this, ensure both values
are asserted to the correct numeric type safely, then compare them using a
tolerance threshold for floating-point numbers instead of direct equality. This
will prevent incorrect difference detection due to minor precision errors.

}
}
if upstreamValue == nil && localValue == nil {
upstreamValue = "same"
}

if localValue == nil && upstreamValue != nil && upstreamValue != "same" {
hasDifference = true
}

upstreamValues[channel.name] = upstreamValue
}

shouldInclude := false

if localValue != nil {
if hasDifference {
shouldInclude = true
}
} else {
if hasUpstreamValue {
shouldInclude = true
}
}

if shouldInclude {
if differences[modelName] == nil {
differences[modelName] = make(map[string]dto.DifferenceItem)
}
differences[modelName][ratioType] = dto.DifferenceItem{
Current: localValue,
Upstreams: upstreamValues,
}
}
}
}

channelHasDiff := make(map[string]bool)
for _, ratioMap := range differences {
for _, item := range ratioMap {
for chName, val := range item.Upstreams {
if val != nil && val != "same" {
channelHasDiff[chName] = true
}
}
}
}

for modelName, ratioMap := range differences {
for ratioType, item := range ratioMap {
for chName := range item.Upstreams {
if !channelHasDiff[chName] {
delete(item.Upstreams, chName)
}
}
differences[modelName][ratioType] = item
}
}

return differences
}

func GetSyncableChannels(c *gin.Context) {
channels, err := model.GetAllChannels(0, 0, true, false)
if err != nil {
c.JSON(http.StatusOK, gin.H{
"success": false,
"message": err.Error(),
})
return
}

var syncableChannels []dto.SyncableChannel
for _, channel := range channels {
if channel.GetBaseURL() != "" {
syncableChannels = append(syncableChannels, dto.SyncableChannel{
ID: channel.Id,
Name: channel.Name,
BaseURL: channel.GetBaseURL(),
Status: channel.Status,
})
}
}

c.JSON(http.StatusOK, gin.H{
"success": true,
"message": "",
"data": syncableChannels,
})
}
Comment on lines +295 to +322

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🛠️ Refactor suggestion

Add input validation and improve error handling.

The GetSyncableChannels function lacks input validation and returns success even when database errors occur.

func GetSyncableChannels(c *gin.Context) {
    channels, err := model.GetAllChannels(0, 0, true, false)
    if err != nil {
-        c.JSON(http.StatusOK, gin.H{
+        c.JSON(http.StatusInternalServerError, gin.H{
            "success": false,
            "message": err.Error(),
        })
        return
    }

    var syncableChannels []dto.SyncableChannel
    for _, channel := range channels {
-        if channel.GetBaseURL() != "" {
+        baseURL := channel.GetBaseURL()
+        if baseURL != "" && isValidURL(baseURL) {  // Reuse the validation function
            syncableChannels = append(syncableChannels, dto.SyncableChannel{
                ID:      channel.Id,
                Name:    channel.Name,
-                BaseURL: channel.GetBaseURL(),
+                BaseURL: baseURL,
                Status:  channel.Status,
            })
        }
    }

    c.JSON(http.StatusOK, gin.H{
        "success": true,
-        "message": "",
+        "message": "Successfully retrieved syncable channels",
        "data":    syncableChannels,
    })
}
🤖 Prompt for AI Agents
In controller/ratio_sync.go around lines 266 to 293, the GetSyncableChannels
function lacks input validation and improperly returns success status even when
database errors occur. Add validation for any input parameters received from the
gin.Context before processing. Modify the error handling to return an
appropriate HTTP error status code (e.g., 500) instead of http.StatusOK when a
database error occurs, and set "success" to false in the JSON response. Ensure
the function only returns success true when no errors happen and valid data is
returned.

Loading