Skip to content
Closed
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
5 changes: 5 additions & 0 deletions common/json.go
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,11 @@ func Marshal(v any) ([]byte, error) {
return json.Marshal(v)
}

// ValidJson 校验数据是否为合法 JSON(替代 encoding/json 的 json.Valid)。
func ValidJson(data []byte) bool {
return json.Valid(data)
}

func GetJsonType(data json.RawMessage) string {
trimmed := bytes.TrimSpace(data)
if len(trimmed) == 0 {
Expand Down
121 changes: 121 additions & 0 deletions common/json_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,8 +2,11 @@ package common

import (
"encoding/json"
"strconv"
"strings"
"testing"

"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)

Expand Down Expand Up @@ -41,3 +44,121 @@ func TestJsonRawMessageToString(t *testing.T) {
})
}
}

// customMarshaler 用于验证底层 JSON 库仍会调用类型自定义的 MarshalJSON/UnmarshalJSON,
// 不依赖 dto 包(避免 common <-> dto 循环引用)。
type customMarshaler struct {
V int
}

func (c customMarshaler) MarshalJSON() ([]byte, error) {
return []byte(`"custom:` + strconv.Itoa(c.V) + `"`), nil
}

func (c *customMarshaler) UnmarshalJSON(b []byte) error {
var s string
if err := Unmarshal(b, &s); err != nil {
return err
}
n, err := strconv.Atoi(strings.TrimPrefix(s, "custom:"))
if err != nil {
return err
}
c.V = n
return nil
}

// TestMarshalStdCompatible 锁定与 encoding/json 字节级一致的关键契约:
// map key 字典序排序(保护依赖 JSON 字节稳定的签名场景)与自定义 Marshaler 仍然生效。
// (HTML 转义与标准库的字节一致性由 TestMarshalMatchesEncodingJSON 覆盖。)
func TestMarshalStdCompatible(t *testing.T) {
tests := []struct {
name string
in any
want string
}{
{"map key 字典序排序", map[string]int{"b": 2, "a": 1, "c": 3}, `{"a":1,"b":2,"c":3}`},
{"嵌套 map 排序", map[string]any{"z": map[string]int{"y": 1, "x": 2}}, `{"z":{"x":2,"y":1}}`},
{"自定义 Marshaler 生效", customMarshaler{42}, `"custom:42"`},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, err := Marshal(tt.in)
require.NoError(t, err)
assert.Equal(t, tt.want, string(got))
})
}
}

// TestMarshalMatchesEncodingJSON 对同一输入直接比对 common.Marshal 与标准库的输出字节,
// 这是「升级底层库零回归」的核心保证。
func TestMarshalMatchesEncodingJSON(t *testing.T) {
inputs := []any{
map[string]int{"b": 2, "a": 1},
map[string]string{"html": `<a href="x">&`},
[]any{1, "two", true, nil},
struct {
Name string `json:"name"`
Tags []string `json:"tags"`
}{"foo", []string{"a", "b"}},
}
for i, in := range inputs {
t.Run(strconv.Itoa(i), func(t *testing.T) {
std, err := json.Marshal(in)
require.NoError(t, err)
got, err := Marshal(in)
require.NoError(t, err)
assert.Equal(t, string(std), string(got))
})
}
}

// TestRoundTrip 验证 RawMessage、json.Number、自定义 Marshaler 与 UnmarshalJsonStr
// 编解码后语义不变。
func TestRoundTrip(t *testing.T) {
t.Run("RawMessage", func(t *testing.T) {
type wrap struct {
Raw json.RawMessage `json:"raw"`
}
src := wrap{Raw: json.RawMessage(`{"k":1}`)}
b, err := Marshal(src)
require.NoError(t, err)
var dst wrap
require.NoError(t, Unmarshal(b, &dst))
assert.JSONEq(t, string(src.Raw), string(dst.Raw))
})

t.Run("json.Number", func(t *testing.T) {
type wrap struct {
N json.Number `json:"n"`
}
b, err := Marshal(wrap{N: "123.456"})
require.NoError(t, err)
assert.Equal(t, `{"n":123.456}`, string(b))
var dst wrap
require.NoError(t, Unmarshal(b, &dst))
assert.Equal(t, json.Number("123.456"), dst.N)
})

t.Run("自定义 Marshaler", func(t *testing.T) {
b, err := Marshal(customMarshaler{7})
require.NoError(t, err)
var c customMarshaler
require.NoError(t, Unmarshal(b, &c))
assert.Equal(t, 7, c.V)
})

t.Run("UnmarshalJsonStr", func(t *testing.T) {
var m map[string]int
require.NoError(t, UnmarshalJsonStr(`{"a":1,"b":2}`, &m))
assert.Equal(t, map[string]int{"a": 1, "b": 2}, m)
})
}

// TestValidJson 覆盖新增的 ValidJson 封装。
func TestValidJson(t *testing.T) {
assert.True(t, ValidJson([]byte(`{"a":1}`)))
assert.True(t, ValidJson([]byte(`[1,2,3]`)))
assert.False(t, ValidJson([]byte(`{"a":}`)))
assert.False(t, ValidJson([]byte(`not json`)))
}
11 changes: 5 additions & 6 deletions common/str.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,6 @@ package common

import (
"encoding/base64"
"encoding/json"
"fmt"
"net/url"
"regexp"
Expand Down Expand Up @@ -46,7 +45,7 @@ func GetRandomString(length int) string {
}

func MapToJsonStr(m map[string]interface{}) string {
bytes, err := json.Marshal(m)
bytes, err := Marshal(m)
if err != nil {
return ""
}
Expand All @@ -64,7 +63,7 @@ func StrToMap(str string) (map[string]interface{}, error) {

func StrToJsonArray(str string) ([]interface{}, error) {
var js []interface{}
err := json.Unmarshal([]byte(str), &js)
err := Unmarshal([]byte(str), &js)
if err != nil {
return nil, err
}
Expand All @@ -73,12 +72,12 @@ func StrToJsonArray(str string) ([]interface{}, error) {

func IsJsonArray(str string) bool {
var js []interface{}
return json.Unmarshal([]byte(str), &js) == nil
return Unmarshal([]byte(str), &js) == nil
}

func IsJsonObject(str string) bool {
var js map[string]interface{}
return json.Unmarshal([]byte(str), &js) == nil
return Unmarshal([]byte(str), &js) == nil
}

func String2Int(str string) int {
Expand Down Expand Up @@ -113,7 +112,7 @@ func GetJsonString(data any) string {
if data == nil {
return ""
}
b, _ := json.Marshal(data)
b, _ := Marshal(data)
return string(b)
}

Expand Down
5 changes: 2 additions & 3 deletions common/topup-ratio.go
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
package common

import (
"encoding/json"
"sync"
)

Expand All @@ -15,7 +14,7 @@ var topupGroupRatioMutex sync.RWMutex
func TopupGroupRatio2JSONString() string {
topupGroupRatioMutex.RLock()
defer topupGroupRatioMutex.RUnlock()
jsonBytes, err := json.Marshal(topupGroupRatio)
jsonBytes, err := Marshal(topupGroupRatio)
if err != nil {
SysError("error marshalling topup group ratio: " + err.Error())
}
Expand All @@ -26,7 +25,7 @@ func UpdateTopupGroupRatioByJSONString(jsonStr string) error {
topupGroupRatioMutex.Lock()
defer topupGroupRatioMutex.Unlock()
topupGroupRatio = make(map[string]float64)
return json.Unmarshal([]byte(jsonStr), &topupGroupRatio)
return Unmarshal([]byte(jsonStr), &topupGroupRatio)
}

func GetTopupGroupRatio(name string) float64 {
Expand Down
5 changes: 2 additions & 3 deletions common/utils.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,6 @@ import (
"crypto/sha256"
"encoding/base64"
"encoding/hex"
"encoding/json"
"fmt"
"html/template"
"io"
Expand Down Expand Up @@ -305,12 +304,12 @@ func GetPointer[T any](v T) *T {

func Any2Type[T any](data any) (T, error) {
var zero T
bytes, err := json.Marshal(data)
bytes, err := Marshal(data)
if err != nil {
return zero, err
}
var res T
err = json.Unmarshal(bytes, &res)
err = Unmarshal(bytes, &res)
if err != nil {
return zero, err
}
Expand Down
23 changes: 11 additions & 12 deletions controller/channel-billing.go
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
package controller

import (
"encoding/json"
"errors"
"fmt"
"io"
Expand Down Expand Up @@ -174,7 +173,7 @@ func updateChannelCloseAIBalance(channel *model.Channel) (float64, error) {
return 0, err
}
response := OpenAICreditGrants{}
err = json.Unmarshal(body, &response)
err = common.Unmarshal(body, &response)
if err != nil {
return 0, err
}
Expand All @@ -189,7 +188,7 @@ func updateChannelOpenAISBBalance(channel *model.Channel) (float64, error) {
return 0, err
}
response := OpenAISBUsageResponse{}
err = json.Unmarshal(body, &response)
err = common.Unmarshal(body, &response)
if err != nil {
return 0, err
}
Expand All @@ -213,7 +212,7 @@ func updateChannelAIProxyBalance(channel *model.Channel) (float64, error) {
return 0, err
}
response := AIProxyUserOverviewResponse{}
err = json.Unmarshal(body, &response)
err = common.Unmarshal(body, &response)
if err != nil {
return 0, err
}
Expand All @@ -232,7 +231,7 @@ func updateChannelAPI2GPTBalance(channel *model.Channel) (float64, error) {
return 0, err
}
response := API2GPTUsageResponse{}
err = json.Unmarshal(body, &response)
err = common.Unmarshal(body, &response)
if err != nil {
return 0, err
}
Expand All @@ -247,7 +246,7 @@ func updateChannelSiliconFlowBalance(channel *model.Channel) (float64, error) {
return 0, err
}
response := SiliconFlowUsageResponse{}
err = json.Unmarshal(body, &response)
err = common.Unmarshal(body, &response)
if err != nil {
return 0, err
}
Expand All @@ -269,7 +268,7 @@ func updateChannelDeepSeekBalance(channel *model.Channel) (float64, error) {
return 0, err
}
response := DeepSeekUsageResponse{}
err = json.Unmarshal(body, &response)
err = common.Unmarshal(body, &response)
if err != nil {
return 0, err
}
Expand Down Expand Up @@ -298,7 +297,7 @@ func updateChannelAIGC2DBalance(channel *model.Channel) (float64, error) {
return 0, err
}
response := APGC2DGPTUsageResponse{}
err = json.Unmarshal(body, &response)
err = common.Unmarshal(body, &response)
if err != nil {
return 0, err
}
Expand All @@ -313,7 +312,7 @@ func updateChannelOpenRouterBalance(channel *model.Channel) (float64, error) {
return 0, err
}
response := OpenRouterCreditResponse{}
err = json.Unmarshal(body, &response)
err = common.Unmarshal(body, &response)
if err != nil {
return 0, err
}
Expand Down Expand Up @@ -343,7 +342,7 @@ func updateChannelMoonshotBalance(channel *model.Channel) (float64, error) {
}

response := MoonshotBalanceResponse{}
err = json.Unmarshal(body, &response)
err = common.Unmarshal(body, &response)
if err != nil {
return 0, err
}
Expand Down Expand Up @@ -396,7 +395,7 @@ func updateChannelBalance(channel *model.Channel) (float64, error) {
return 0, err
}
subscription := OpenAISubscriptionResponse{}
err = json.Unmarshal(body, &subscription)
err = common.Unmarshal(body, &subscription)
if err != nil {
return 0, err
}
Expand All @@ -412,7 +411,7 @@ func updateChannelBalance(channel *model.Channel) (float64, error) {
return 0, err
}
usage := OpenAIUsageResponse{}
err = json.Unmarshal(body, &usage)
err = common.Unmarshal(body, &usage)
if err != nil {
return 0, err
}
Expand Down
Loading