+
+ {t('利润比例')}
+ {
+ setAistarsLabProfitPercent(value ?? 0);
+ setAistarsLabResult(null);
+ }}
+ className='w-full sm:w-28'
+ disabled={aistarsLabLoading}
+ />
+ %
+
}
className='w-full md:w-auto'
@@ -639,6 +664,12 @@ export default function UpstreamRatioSync(props) {
{t('映射')}: {aistarsLabResult.mapping_changes?.length || 0}
+
+ {t('利润比例')}:{' '}
+ {Number.isFinite(aistarsLabResult.markup_rate)
+ ? `${Math.round((aistarsLabResult.markup_rate - 1) * 10000) / 100}%`
+ : '-'}
+
Date: Sat, 4 Jul 2026 23:52:10 +0800
Subject: [PATCH 13/24] fix sora task string error parsing
---
relay/channel/task/sora/adaptor.go | 59 ++++++++++++++++++-------
relay/channel/task/sora/adaptor_test.go | 39 ++++++++++++++++
2 files changed, 81 insertions(+), 17 deletions(-)
create mode 100644 relay/channel/task/sora/adaptor_test.go
diff --git a/relay/channel/task/sora/adaptor.go b/relay/channel/task/sora/adaptor.go
index e9029aa20d46..4a22ed134877 100644
--- a/relay/channel/task/sora/adaptor.go
+++ b/relay/channel/task/sora/adaptor.go
@@ -39,22 +39,47 @@ type ImageURL struct {
}
type responseTask struct {
- ID string `json:"id"`
- TaskID string `json:"task_id,omitempty"` //兼容旧接口
- Object string `json:"object"`
- Model string `json:"model"`
- Status string `json:"status"`
- Progress int `json:"progress"`
- CreatedAt int64 `json:"created_at"`
- CompletedAt int64 `json:"completed_at,omitempty"`
- ExpiresAt int64 `json:"expires_at,omitempty"`
- Seconds string `json:"seconds,omitempty"`
- Size string `json:"size,omitempty"`
- RemixedFromVideoID string `json:"remixed_from_video_id,omitempty"`
- Error *struct {
- Message string `json:"message"`
- Code string `json:"code"`
- } `json:"error,omitempty"`
+ ID string `json:"id"`
+ TaskID string `json:"task_id,omitempty"` //兼容旧接口
+ Object string `json:"object"`
+ Model string `json:"model"`
+ Status string `json:"status"`
+ Progress int `json:"progress"`
+ CreatedAt int64 `json:"created_at"`
+ CompletedAt int64 `json:"completed_at,omitempty"`
+ ExpiresAt int64 `json:"expires_at,omitempty"`
+ Seconds string `json:"seconds,omitempty"`
+ Size string `json:"size,omitempty"`
+ RemixedFromVideoID string `json:"remixed_from_video_id,omitempty"`
+ Error *responseTaskError `json:"error,omitempty"`
+}
+
+type responseTaskError struct {
+ Message string `json:"message"`
+ Code string `json:"code"`
+}
+
+func (e *responseTaskError) UnmarshalJSON(data []byte) error {
+ trimmed := strings.TrimSpace(string(data))
+ if trimmed == "" || trimmed == "null" {
+ return nil
+ }
+ if strings.HasPrefix(trimmed, `"`) {
+ var message string
+ if err := common.Unmarshal(data, &message); err != nil {
+ return err
+ }
+ e.Message = message
+ return nil
+ }
+
+ type alias responseTaskError
+ var parsed alias
+ if err := common.Unmarshal(data, &parsed); err != nil {
+ return err
+ }
+ *e = responseTaskError(parsed)
+ return nil
}
// ============================
@@ -307,7 +332,7 @@ func (a *TaskAdaptor) ParseTaskResult(respBody []byte) (*relaycommon.TaskInfo, e
// Url intentionally left empty — the caller constructs the proxy URL using the public task ID
case "failed", "cancelled":
taskResult.Status = model.TaskStatusFailure
- if resTask.Error != nil {
+ if resTask.Error != nil && resTask.Error.Message != "" {
taskResult.Reason = resTask.Error.Message
} else {
taskResult.Reason = "task failed"
diff --git a/relay/channel/task/sora/adaptor_test.go b/relay/channel/task/sora/adaptor_test.go
new file mode 100644
index 000000000000..d625405cf64d
--- /dev/null
+++ b/relay/channel/task/sora/adaptor_test.go
@@ -0,0 +1,39 @@
+package sora
+
+import (
+ "testing"
+
+ "github.com/QuantumNous/new-api/model"
+ "github.com/stretchr/testify/assert"
+ "github.com/stretchr/testify/require"
+)
+
+func TestParseTaskResultFailedWithStringError(t *testing.T) {
+ adaptor := &TaskAdaptor{}
+
+ taskInfo, err := adaptor.ParseTaskResult([]byte(`{
+ "id": "task_upstream",
+ "status": "failed",
+ "error": "safety system rejected this request"
+ }`))
+
+ require.NoError(t, err)
+ require.NotNil(t, taskInfo)
+ assert.Equal(t, model.TaskStatusFailure, taskInfo.Status)
+ assert.Equal(t, "safety system rejected this request", taskInfo.Reason)
+}
+
+func TestParseTaskResultFailedWithObjectError(t *testing.T) {
+ adaptor := &TaskAdaptor{}
+
+ taskInfo, err := adaptor.ParseTaskResult([]byte(`{
+ "id": "task_upstream",
+ "status": "failed",
+ "error": {"message": "invalid prompt", "code": "invalid_request"}
+ }`))
+
+ require.NoError(t, err)
+ require.NotNil(t, taskInfo)
+ assert.Equal(t, model.TaskStatusFailure, taskInfo.Status)
+ assert.Equal(t, "invalid prompt", taskInfo.Reason)
+}
From 99384cb46d9d4eaf782b20d911a849d9cf1a4d0c Mon Sep 17 00:00:00 2001
From: "649985538@qq.com" <649985538@qq.com>
Date: Sun, 5 Jul 2026 00:01:11 +0800
Subject: [PATCH 14/24] fix sora task error responses without status
---
relay/channel/task/sora/adaptor.go | 4 ++++
relay/channel/task/sora/adaptor_test.go | 17 +++++++++++++++++
2 files changed, 21 insertions(+)
diff --git a/relay/channel/task/sora/adaptor.go b/relay/channel/task/sora/adaptor.go
index 4a22ed134877..18643f9cf7a0 100644
--- a/relay/channel/task/sora/adaptor.go
+++ b/relay/channel/task/sora/adaptor.go
@@ -338,6 +338,10 @@ func (a *TaskAdaptor) ParseTaskResult(respBody []byte) (*relaycommon.TaskInfo, e
taskResult.Reason = "task failed"
}
default:
+ if resTask.Error != nil && resTask.Error.Message != "" {
+ taskResult.Status = model.TaskStatusFailure
+ taskResult.Reason = resTask.Error.Message
+ }
}
if resTask.Progress > 0 && resTask.Progress < 100 {
taskResult.Progress = fmt.Sprintf("%d%%", resTask.Progress)
diff --git a/relay/channel/task/sora/adaptor_test.go b/relay/channel/task/sora/adaptor_test.go
index d625405cf64d..3199a7cf28f4 100644
--- a/relay/channel/task/sora/adaptor_test.go
+++ b/relay/channel/task/sora/adaptor_test.go
@@ -37,3 +37,20 @@ func TestParseTaskResultFailedWithObjectError(t *testing.T) {
assert.Equal(t, model.TaskStatusFailure, taskInfo.Status)
assert.Equal(t, "invalid prompt", taskInfo.Reason)
}
+
+func TestParseTaskResultErrorWithoutStatus(t *testing.T) {
+ adaptor := &TaskAdaptor{}
+
+ taskInfo, err := adaptor.ParseTaskResult([]byte(`{
+ "code": "Client specified an invalid argument",
+ "error": "Generated video rejected by content moderation.",
+ "id": "task_upstream",
+ "task_id": "task_upstream",
+ "model": "grok-image-video"
+ }`))
+
+ require.NoError(t, err)
+ require.NotNil(t, taskInfo)
+ assert.Equal(t, model.TaskStatusFailure, taskInfo.Status)
+ assert.Equal(t, "Generated video rejected by content moderation.", taskInfo.Reason)
+}
From b1cd944bdfb2982c3d7b16c1447ec7489424683a Mon Sep 17 00:00:00 2001
From: "649985538@qq.com" <649985538@qq.com>
Date: Sun, 5 Jul 2026 00:48:27 +0800
Subject: [PATCH 15/24] fix pricing completion ratio override
---
setting/ratio_setting/model_ratio.go | 27 +++---------
setting/ratio_setting/model_ratio_test.go | 53 +++++++++++++++++++++++
2 files changed, 60 insertions(+), 20 deletions(-)
create mode 100644 setting/ratio_setting/model_ratio_test.go
diff --git a/setting/ratio_setting/model_ratio.go b/setting/ratio_setting/model_ratio.go
index 7556fd9482c7..3a4a0768cf2a 100644
--- a/setting/ratio_setting/model_ratio.go
+++ b/setting/ratio_setting/model_ratio.go
@@ -443,18 +443,14 @@ func UpdateCompletionRatioByJSONString(jsonStr string) error {
func GetCompletionRatio(name string) float64 {
name = FormatMatchingModelName(name)
- if strings.Contains(name, "/") {
- if ratio, ok := completionRatioMap.Get(name); ok {
- return ratio
- }
+ if ratio, ok := completionRatioMap.Get(name); ok {
+ return ratio
}
+
hardCodedRatio, contain := getHardcodedCompletionModelRatio(name)
if contain {
return hardCodedRatio
}
- if ratio, ok := completionRatioMap.Get(name); ok {
- return ratio
- }
return hardCodedRatio
}
@@ -466,12 +462,10 @@ type CompletionRatioInfo struct {
func GetCompletionRatioInfo(name string) CompletionRatioInfo {
name = FormatMatchingModelName(name)
- if strings.Contains(name, "/") {
- if ratio, ok := completionRatioMap.Get(name); ok {
- return CompletionRatioInfo{
- Ratio: ratio,
- Locked: false,
- }
+ if ratio, ok := completionRatioMap.Get(name); ok {
+ return CompletionRatioInfo{
+ Ratio: ratio,
+ Locked: false,
}
}
@@ -483,13 +477,6 @@ func GetCompletionRatioInfo(name string) CompletionRatioInfo {
}
}
- if ratio, ok := completionRatioMap.Get(name); ok {
- return CompletionRatioInfo{
- Ratio: ratio,
- Locked: false,
- }
- }
-
return CompletionRatioInfo{
Ratio: hardCodedRatio,
Locked: false,
diff --git a/setting/ratio_setting/model_ratio_test.go b/setting/ratio_setting/model_ratio_test.go
new file mode 100644
index 000000000000..38a5530b4947
--- /dev/null
+++ b/setting/ratio_setting/model_ratio_test.go
@@ -0,0 +1,53 @@
+package ratio_setting
+
+import "testing"
+
+func TestConfiguredCompletionRatioOverridesHardcodedGPT5Ratio(t *testing.T) {
+ originalCompletionRatio := CompletionRatio2JSONString()
+ t.Cleanup(func() {
+ if err := UpdateCompletionRatioByJSONString(originalCompletionRatio); err != nil {
+ t.Fatalf("restore completion ratio: %v", err)
+ }
+ })
+
+ if err := UpdateCompletionRatioByJSONString(`{"gpt-5.5":6}`); err != nil {
+ t.Fatalf("update completion ratio: %v", err)
+ }
+
+ if got := GetCompletionRatio("gpt-5.5"); got != 6 {
+ t.Fatalf("GetCompletionRatio() = %v, want 6", got)
+ }
+
+ info := GetCompletionRatioInfo("gpt-5.5")
+ if info.Ratio != 6 {
+ t.Fatalf("GetCompletionRatioInfo().Ratio = %v, want 6", info.Ratio)
+ }
+ if info.Locked {
+ t.Fatal("GetCompletionRatioInfo().Locked = true, want false for configured ratio")
+ }
+}
+
+func TestHardcodedCompletionRatioAppliesWhenGPT5RatioIsNotConfigured(t *testing.T) {
+ originalCompletionRatio := CompletionRatio2JSONString()
+ t.Cleanup(func() {
+ if err := UpdateCompletionRatioByJSONString(originalCompletionRatio); err != nil {
+ t.Fatalf("restore completion ratio: %v", err)
+ }
+ })
+
+ if err := UpdateCompletionRatioByJSONString(`{}`); err != nil {
+ t.Fatalf("update completion ratio: %v", err)
+ }
+
+ if got := GetCompletionRatio("gpt-5.5"); got != 8 {
+ t.Fatalf("GetCompletionRatio() = %v, want 8", got)
+ }
+
+ info := GetCompletionRatioInfo("gpt-5.5")
+ if info.Ratio != 8 {
+ t.Fatalf("GetCompletionRatioInfo().Ratio = %v, want 8", info.Ratio)
+ }
+ if !info.Locked {
+ t.Fatal("GetCompletionRatioInfo().Locked = false, want true for hardcoded ratio")
+ }
+}
From a3fe3df2d8bad410da8ed7fa90a2803ef2396b2d Mon Sep 17 00:00:00 2001
From: "649985538@qq.com" <649985538@qq.com>
Date: Tue, 7 Jul 2026 00:49:34 +0800
Subject: [PATCH 16/24] fix topup return url override
---
controller/topup.go | 11 ++++++++++-
1 file changed, 10 insertions(+), 1 deletion(-)
diff --git a/controller/topup.go b/controller/topup.go
index 86d361a349cb..4ef20ac44ea0 100644
--- a/controller/topup.go
+++ b/controller/topup.go
@@ -117,6 +117,7 @@ func GetTopUpInfo(c *gin.Context) {
type EpayRequest struct {
Amount int64 `json:"amount"`
PaymentMethod string `json:"payment_method"`
+ ReturnUrl string `json:"return_url,omitempty"`
}
type AmountRequest struct {
@@ -218,7 +219,15 @@ func RequestEpay(c *gin.Context) {
}
callBackAddress := service.GetCallbackAddress()
- returnUrl, _ := url.Parse(system_setting.ServerAddress + "/console/log")
+ returnUrlValue := system_setting.ServerAddress + "/console/log"
+ if req.ReturnUrl != "" {
+ if err := common.ValidateRedirectURL(req.ReturnUrl); err != nil {
+ c.JSON(http.StatusBadRequest, gin.H{"message": "支付完成重定向URL不在可信任域名列表中", "data": ""})
+ return
+ }
+ returnUrlValue = req.ReturnUrl
+ }
+ returnUrl, _ := url.Parse(returnUrlValue)
notifyUrl, _ := url.Parse(callBackAddress + "/api/user/epay/notify")
tradeNo := fmt.Sprintf("%s%d", common.GetRandomString(6), time.Now().Unix())
tradeNo = fmt.Sprintf("USR%dNO%s", id, tradeNo)
From e63e970806c7fca8b1d7a46cc5cf29a352a75896 Mon Sep 17 00:00:00 2001
From: "649985538@qq.com" <649985538@qq.com>
Date: Wed, 8 Jul 2026 16:13:09 +0800
Subject: [PATCH 17/24] skip task retry on forbidden upstream responses
---
controller/relay.go | 3 +++
controller/relay_retry_test.go | 19 +++++++++++++++++++
2 files changed, 22 insertions(+)
create mode 100644 controller/relay_retry_test.go
diff --git a/controller/relay.go b/controller/relay.go
index a59b3abd66ec..b3b9dfce386e 100644
--- a/controller/relay.go
+++ b/controller/relay.go
@@ -634,6 +634,9 @@ func shouldRetryTaskRelay(c *gin.Context, channelId int, taskErr *dto.TaskError,
if taskErr.StatusCode == http.StatusBadRequest {
return false
}
+ if taskErr.StatusCode == http.StatusForbidden {
+ return false
+ }
if taskErr.StatusCode == 408 {
// azure处理超时不重试
return false
diff --git a/controller/relay_retry_test.go b/controller/relay_retry_test.go
new file mode 100644
index 000000000000..00a8de668c97
--- /dev/null
+++ b/controller/relay_retry_test.go
@@ -0,0 +1,19 @@
+package controller
+
+import (
+ "net/http"
+ "testing"
+
+ "github.com/QuantumNous/new-api/dto"
+ "github.com/gin-gonic/gin"
+ "github.com/stretchr/testify/require"
+)
+
+func TestShouldRetryTaskRelaySkipsForbidden(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ ctx, _ := gin.CreateTestContext(nil)
+
+ retry := shouldRetryTaskRelay(ctx, 19, &dto.TaskError{StatusCode: http.StatusForbidden}, 5)
+
+ require.False(t, retry)
+}
From d99bfbce62470b15a22c187365e7eaf84b25b0fa Mon Sep 17 00:00:00 2001
From: "649985538@qq.com" <649985538@qq.com>
Date: Thu, 9 Jul 2026 22:43:27 +0800
Subject: [PATCH 18/24] add built-in image resolution pricing
---
dto/openai_image.go | 128 +++++++++++++++++++++++++++++++++++++
dto/openai_image_test.go | 77 ++++++++++++++++++++++
relay/helper/price.go | 12 +++-
relay/helper/price_test.go | 54 ++++++++++++++++
service/text_quota_test.go | 30 +++++++++
types/request_meta.go | 1 +
6 files changed, 301 insertions(+), 1 deletion(-)
create mode 100644 dto/openai_image_test.go
create mode 100644 relay/helper/price_test.go
diff --git a/dto/openai_image.go b/dto/openai_image.go
index 52986fbfd59d..ccb5105a614d 100644
--- a/dto/openai_image.go
+++ b/dto/openai_image.go
@@ -2,6 +2,7 @@ package dto
import (
"encoding/json"
+ "math"
"reflect"
"strings"
@@ -124,9 +125,135 @@ func indexComma(s string) int {
return -1
}
+func normalizeImageQuality(quality string) string {
+ switch strings.ToLower(strings.TrimSpace(quality)) {
+ case "low", "medium", "high":
+ return strings.ToLower(strings.TrimSpace(quality))
+ default:
+ return "medium"
+ }
+}
+
+func parseImageSize(size string) (int, int, bool) {
+ size = strings.ToLower(strings.TrimSpace(size))
+ if size == "" || size == "auto" {
+ size = "1024x1024"
+ }
+ parts := strings.Split(strings.ToLower(strings.TrimSpace(size)), "x")
+ if len(parts) != 2 {
+ return 0, 0, false
+ }
+ width := common.String2Int(strings.TrimSpace(parts[0]))
+ height := common.String2Int(strings.TrimSpace(parts[1]))
+ if width <= 0 || height <= 0 {
+ return 0, 0, false
+ }
+ return width, height, true
+}
+
+func imageSizeTier(size string) (string, bool) {
+ width, height, ok := parseImageSize(size)
+ if !ok {
+ return "", false
+ }
+ longEdge := width
+ if height > longEdge {
+ longEdge = height
+ }
+ switch {
+ case longEdge <= 1024:
+ return "1k", true
+ case longEdge <= 2048:
+ return "2k", true
+ case longEdge <= 4096:
+ return "4k", true
+ default:
+ return "", false
+ }
+}
+
+func gptImage2UnitPrice(size string, quality string) (float64, bool) {
+ width, height, ok := parseImageSize(size)
+ if !ok {
+ return 0, false
+ }
+ if width%16 != 0 || height%16 != 0 {
+ return 0, false
+ }
+ pixels := width * height
+ if pixels < 655360 || pixels > 8294400 {
+ return 0, false
+ }
+ longEdge := width
+ shortEdge := height
+ if height > width {
+ longEdge = height
+ shortEdge = width
+ }
+ if longEdge > 3840 || float64(longEdge)/float64(shortEdge) > 3 {
+ return 0, false
+ }
+
+ qualityGrid := map[string]int{
+ "low": 16,
+ "medium": 48,
+ "high": 96,
+ }[normalizeImageQuality(quality)]
+ shortGrid := int(math.Round(float64(qualityGrid) * float64(shortEdge) / float64(longEdge)))
+ widthGrid := shortGrid
+ heightGrid := qualityGrid
+ if width >= height {
+ widthGrid = qualityGrid
+ heightGrid = shortGrid
+ }
+ outputTokens := math.Ceil(float64(widthGrid*heightGrid) * float64(2000000+pixels) / 4000000)
+ return outputTokens * 30 / 1000000, true
+}
+
+func builtInImageUnitPrice(model string, size string, quality string) (float64, bool) {
+ model = strings.ToLower(strings.TrimSpace(model))
+ if model == "gpt-image-2" {
+ return gptImage2UnitPrice(size, quality)
+ }
+
+ tier, ok := imageSizeTier(size)
+ if !ok {
+ return 0, false
+ }
+
+ switch model {
+ case "gemini-3.1-flash-image", "nano-banana-2":
+ switch tier {
+ case "1k":
+ return 0.067, true
+ case "2k":
+ return 0.101, true
+ case "4k":
+ return 0.151, true
+ }
+ case "gemini-3-pro-image", "nano-banana-pro":
+ switch tier {
+ case "1k", "2k":
+ return 0.134, true
+ case "4k":
+ return 0.240, true
+ }
+ case "gemini-2.5-flash-image", "nano-banana":
+ if tier == "1k" {
+ return 0.039, true
+ }
+ case "gemini-3.1-flash-lite-image":
+ if tier == "1k" {
+ return 0.0336, true
+ }
+ }
+ return 0, false
+}
+
func (i *ImageRequest) GetTokenCountMeta() *types.TokenCountMeta {
var sizeRatio = 1.0
var qualityRatio = 1.0
+ imageUnitPrice, _ := builtInImageUnitPrice(i.Model, i.Size, i.Quality)
if strings.HasPrefix(i.Model, "dall-e") {
// Size
@@ -156,6 +283,7 @@ func (i *ImageRequest) GetTokenCountMeta() *types.TokenCountMeta {
CombineText: i.Prompt,
MaxTokens: 1584,
ImagePriceRatio: sizeRatio * qualityRatio,
+ ImageUnitPrice: imageUnitPrice,
}
}
diff --git a/dto/openai_image_test.go b/dto/openai_image_test.go
new file mode 100644
index 000000000000..66de219378a3
--- /dev/null
+++ b/dto/openai_image_test.go
@@ -0,0 +1,77 @@
+package dto
+
+import (
+ "testing"
+
+ "github.com/stretchr/testify/require"
+)
+
+func TestImageRequestBuiltInUnitPrice(t *testing.T) {
+ tests := []struct {
+ name string
+ request ImageRequest
+ want float64
+ }{
+ {
+ name: "gpt-image-2 medium 2k square",
+ request: ImageRequest{
+ Model: "gpt-image-2",
+ Size: "2048x2048",
+ Quality: "medium",
+ },
+ want: 0.10704,
+ },
+ {
+ name: "gpt-image-2 high 4k landscape",
+ request: ImageRequest{
+ Model: "gpt-image-2",
+ Size: "3840x2160",
+ Quality: "high",
+ },
+ want: 0.40026,
+ },
+ {
+ name: "banana 2 4k",
+ request: ImageRequest{
+ Model: "gemini-3.1-flash-image",
+ Size: "4096x4096",
+ },
+ want: 0.151,
+ },
+ {
+ name: "banana pro 2k",
+ request: ImageRequest{
+ Model: "gemini-3-pro-image",
+ Size: "2048x2048",
+ },
+ want: 0.134,
+ },
+ {
+ name: "empty size defaults to 1k",
+ request: ImageRequest{
+ Model: "gemini-3.1-flash-image",
+ },
+ want: 0.067,
+ },
+ }
+
+ for _, test := range tests {
+ t.Run(test.name, func(t *testing.T) {
+ meta := test.request.GetTokenCountMeta()
+ require.InDelta(t, test.want, meta.ImageUnitPrice, 0.000001)
+ })
+ }
+}
+
+func TestImageRequestUnknownBuiltInPriceKeepsLegacyImageRatio(t *testing.T) {
+ req := ImageRequest{
+ Model: "dall-e-3",
+ Size: "1024x1792",
+ Quality: "hd",
+ }
+
+ meta := req.GetTokenCountMeta()
+
+ require.Zero(t, meta.ImageUnitPrice)
+ require.Equal(t, 3.0, meta.ImagePriceRatio)
+}
diff --git a/relay/helper/price.go b/relay/helper/price.go
index 8ba0ee8f0844..748a6254518f 100644
--- a/relay/helper/price.go
+++ b/relay/helper/price.go
@@ -62,7 +62,14 @@ func HandleGroupRatio(ctx *gin.Context, relayInfo *relaycommon.RelayInfo) types.
}
func ModelPriceHelper(c *gin.Context, info *relaycommon.RelayInfo, promptTokens int, meta *types.TokenCountMeta) (types.PriceData, error) {
+ if meta == nil {
+ meta = &types.TokenCountMeta{}
+ }
modelPrice, usePrice := ratio_setting.GetModelPrice(info.OriginModelName, false)
+ if meta.ImageUnitPrice > 0 {
+ modelPrice = meta.ImageUnitPrice
+ usePrice = true
+ }
groupRatioInfo := HandleGroupRatio(c, info)
@@ -106,7 +113,10 @@ func ModelPriceHelper(c *gin.Context, info *relaycommon.RelayInfo, promptTokens
ratio := modelRatio * groupRatioInfo.GroupRatio
preConsumedQuota = int(float64(preConsumedTokens) * ratio)
} else {
- if meta.ImagePriceRatio != 0 {
+ if meta.ImageUnitPrice > 0 {
+ // Built-in image prices already represent the final single-image price
+ // for the requested model/size/quality.
+ } else if meta.ImagePriceRatio != 0 {
modelPrice = modelPrice * meta.ImagePriceRatio
}
preConsumedQuota = int(modelPrice * common.QuotaPerUnit * groupRatioInfo.GroupRatio)
diff --git a/relay/helper/price_test.go b/relay/helper/price_test.go
new file mode 100644
index 000000000000..c4a79a3e21f5
--- /dev/null
+++ b/relay/helper/price_test.go
@@ -0,0 +1,54 @@
+package helper
+
+import (
+ "net/http/httptest"
+ "testing"
+
+ "github.com/QuantumNous/new-api/common"
+ relaycommon "github.com/QuantumNous/new-api/relay/common"
+ "github.com/QuantumNous/new-api/types"
+
+ "github.com/gin-gonic/gin"
+ "github.com/stretchr/testify/require"
+)
+
+func TestModelPriceHelperUsesBuiltInImageUnitPrice(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ w := httptest.NewRecorder()
+ ctx, _ := gin.CreateTestContext(w)
+
+ info := &relaycommon.RelayInfo{
+ OriginModelName: "gemini-3.1-flash-image",
+ UsingGroup: "default",
+ }
+
+ priceData, err := ModelPriceHelper(ctx, info, 1, &types.TokenCountMeta{
+ ImageUnitPrice: 0.101,
+ })
+
+ require.NoError(t, err)
+ require.True(t, priceData.UsePrice)
+ require.Equal(t, 0.101, priceData.ModelPrice)
+ require.Equal(t, int(0.101*common.QuotaPerUnit), priceData.QuotaToPreConsume)
+}
+
+func TestModelPriceHelperBuiltInImageUnitPriceSkipsImageRatio(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ w := httptest.NewRecorder()
+ ctx, _ := gin.CreateTestContext(w)
+
+ info := &relaycommon.RelayInfo{
+ OriginModelName: "gpt-image-2",
+ UsingGroup: "default",
+ }
+
+ priceData, err := ModelPriceHelper(ctx, info, 1, &types.TokenCountMeta{
+ ImageUnitPrice: 0.10704,
+ ImagePriceRatio: 16,
+ })
+
+ require.NoError(t, err)
+ require.True(t, priceData.UsePrice)
+ require.Equal(t, 0.10704, priceData.ModelPrice)
+ require.Equal(t, int(0.10704*common.QuotaPerUnit), priceData.QuotaToPreConsume)
+}
diff --git a/service/text_quota_test.go b/service/text_quota_test.go
index e995de17ae8b..fbfe80d46ea5 100644
--- a/service/text_quota_test.go
+++ b/service/text_quota_test.go
@@ -316,3 +316,33 @@ func TestCalculateTextQuotaSummaryKeepsPrePRClaudeOpenRouterBilling(t *testing.T
require.Equal(t, 172, summary.PromptTokens)
require.Equal(t, 798, summary.Quota)
}
+
+func TestCalculateTextQuotaSummaryAppliesImageUnitPriceWithCountAndGroupRatio(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ w := httptest.NewRecorder()
+ ctx, _ := gin.CreateTestContext(w)
+
+ relayInfo := &relaycommon.RelayInfo{
+ OriginModelName: "gemini-3.1-flash-image",
+ PriceData: types.PriceData{
+ UsePrice: true,
+ ModelPrice: 0.101,
+ OtherRatios: map[string]float64{
+ "n": 3,
+ },
+ GroupRatioInfo: types.GroupRatioInfo{
+ GroupRatio: 1.5,
+ },
+ },
+ StartTime: time.Now(),
+ }
+
+ usage := &dto.Usage{
+ PromptTokens: 1,
+ TotalTokens: 1,
+ }
+
+ summary := calculateTextQuotaSummary(ctx, relayInfo, usage)
+
+ require.Equal(t, 227250, summary.Quota)
+}
diff --git a/types/request_meta.go b/types/request_meta.go
index 476ea0524dfb..b30678bff1a7 100644
--- a/types/request_meta.go
+++ b/types/request_meta.go
@@ -27,6 +27,7 @@ type TokenCountMeta struct {
MaxTokens int `json:"max_tokens,omitempty"` // Maximum tokens allowed in the request
ImagePriceRatio float64 `json:"image_ratio,omitempty"` // Ratio for image size, if applicable
+ ImageUnitPrice float64 `json:"image_unit_price,omitempty"`
//IsStreaming bool `json:"is_streaming,omitempty"` // Indicates if the request is streaming
}
From 2f618a0dc92f34741de4666f5453f4f0732bef9f Mon Sep 17 00:00:00 2001
From: "649985538@qq.com" <649985538@qq.com>
Date: Thu, 9 Jul 2026 23:05:10 +0800
Subject: [PATCH 19/24] price image group by resolution
---
dto/openai_image.go | 26 ++++++++++++++++++++++----
dto/openai_image_test.go | 24 ++++++++++++++++++++++++
relay/helper/price.go | 20 ++++++++++++++++----
relay/helper/price_test.go | 37 +++++++++++++++++++++++++++++++++++++
types/request_meta.go | 5 +++--
5 files changed, 102 insertions(+), 10 deletions(-)
diff --git a/dto/openai_image.go b/dto/openai_image.go
index ccb5105a614d..03883ba027f9 100644
--- a/dto/openai_image.go
+++ b/dto/openai_image.go
@@ -172,6 +172,22 @@ func imageSizeTier(size string) (string, bool) {
}
}
+func imageGroupUnitPrice(size string) (float64, bool) {
+ tier, ok := imageSizeTier(size)
+ if !ok {
+ return 0, false
+ }
+ switch tier {
+ case "1k":
+ return 0.10, true
+ case "2k":
+ return 0.14, true
+ case "4k":
+ return 0.20, true
+ }
+ return 0, false
+}
+
func gptImage2UnitPrice(size string, quality string) (float64, bool) {
width, height, ok := parseImageSize(size)
if !ok {
@@ -254,6 +270,7 @@ func (i *ImageRequest) GetTokenCountMeta() *types.TokenCountMeta {
var sizeRatio = 1.0
var qualityRatio = 1.0
imageUnitPrice, _ := builtInImageUnitPrice(i.Model, i.Size, i.Quality)
+ imageGroupUnitPrice, _ := imageGroupUnitPrice(i.Size)
if strings.HasPrefix(i.Model, "dall-e") {
// Size
@@ -280,10 +297,11 @@ func (i *ImageRequest) GetTokenCountMeta() *types.TokenCountMeta {
// Including n here caused double-counting for channels that also
// set OtherRatio("n") (e.g. Ali/Bailian).
return &types.TokenCountMeta{
- CombineText: i.Prompt,
- MaxTokens: 1584,
- ImagePriceRatio: sizeRatio * qualityRatio,
- ImageUnitPrice: imageUnitPrice,
+ CombineText: i.Prompt,
+ MaxTokens: 1584,
+ ImagePriceRatio: sizeRatio * qualityRatio,
+ ImageUnitPrice: imageUnitPrice,
+ ImageGroupUnitPrice: imageGroupUnitPrice,
}
}
diff --git a/dto/openai_image_test.go b/dto/openai_image_test.go
index 66de219378a3..85c3f908e9b0 100644
--- a/dto/openai_image_test.go
+++ b/dto/openai_image_test.go
@@ -63,6 +63,30 @@ func TestImageRequestBuiltInUnitPrice(t *testing.T) {
}
}
+func TestImageRequestImageGroupUnitPrice(t *testing.T) {
+ tests := []struct {
+ name string
+ size string
+ want float64
+ }{
+ {name: "empty defaults to 1k", want: 0.10},
+ {name: "1k", size: "1024x1024", want: 0.10},
+ {name: "2k", size: "2048x2048", want: 0.14},
+ {name: "4k", size: "4096x4096", want: 0.20},
+ }
+
+ for _, test := range tests {
+ t.Run(test.name, func(t *testing.T) {
+ req := ImageRequest{
+ Model: "any-image-model",
+ Size: test.size,
+ }
+ meta := req.GetTokenCountMeta()
+ require.InDelta(t, test.want, meta.ImageGroupUnitPrice, 0.000001)
+ })
+ }
+}
+
func TestImageRequestUnknownBuiltInPriceKeepsLegacyImageRatio(t *testing.T) {
req := ImageRequest{
Model: "dall-e-3",
diff --git a/relay/helper/price.go b/relay/helper/price.go
index 748a6254518f..eb1c50ad434b 100644
--- a/relay/helper/price.go
+++ b/relay/helper/price.go
@@ -2,6 +2,7 @@ package helper
import (
"fmt"
+ "strings"
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/logger"
@@ -14,6 +15,10 @@ import (
"github.com/gin-gonic/gin"
)
+func isImagePricingGroup(group string) bool {
+ return strings.EqualFold(strings.TrimSpace(group), "image")
+}
+
func modelPriceNotConfiguredError(modelName string, userId int) error {
if model.IsAdmin(userId) {
return fmt.Errorf(
@@ -66,13 +71,20 @@ func ModelPriceHelper(c *gin.Context, info *relaycommon.RelayInfo, promptTokens
meta = &types.TokenCountMeta{}
}
modelPrice, usePrice := ratio_setting.GetModelPrice(info.OriginModelName, false)
- if meta.ImageUnitPrice > 0 {
+
+ groupRatioInfo := HandleGroupRatio(c, info)
+
+ imageUnitPriceOverride := false
+ if isImagePricingGroup(info.UsingGroup) && meta.ImageGroupUnitPrice > 0 {
+ modelPrice = meta.ImageGroupUnitPrice
+ usePrice = true
+ imageUnitPriceOverride = true
+ } else if meta.ImageUnitPrice > 0 {
modelPrice = meta.ImageUnitPrice
usePrice = true
+ imageUnitPriceOverride = true
}
- groupRatioInfo := HandleGroupRatio(c, info)
-
var preConsumedQuota int
var modelRatio float64
var completionRatio float64
@@ -113,7 +125,7 @@ func ModelPriceHelper(c *gin.Context, info *relaycommon.RelayInfo, promptTokens
ratio := modelRatio * groupRatioInfo.GroupRatio
preConsumedQuota = int(float64(preConsumedTokens) * ratio)
} else {
- if meta.ImageUnitPrice > 0 {
+ if imageUnitPriceOverride {
// Built-in image prices already represent the final single-image price
// for the requested model/size/quality.
} else if meta.ImagePriceRatio != 0 {
diff --git a/relay/helper/price_test.go b/relay/helper/price_test.go
index c4a79a3e21f5..0a5d69f13be6 100644
--- a/relay/helper/price_test.go
+++ b/relay/helper/price_test.go
@@ -52,3 +52,40 @@ func TestModelPriceHelperBuiltInImageUnitPriceSkipsImageRatio(t *testing.T) {
require.Equal(t, 0.10704, priceData.ModelPrice)
require.Equal(t, int(0.10704*common.QuotaPerUnit), priceData.QuotaToPreConsume)
}
+
+func TestModelPriceHelperImageGroupUsesResolutionUnitPriceForAnyModel(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ w := httptest.NewRecorder()
+ ctx, _ := gin.CreateTestContext(w)
+
+ info := &relaycommon.RelayInfo{
+ OriginModelName: "unknown-image-model",
+ UsingGroup: "image",
+ }
+
+ priceData, err := ModelPriceHelper(ctx, info, 1, &types.TokenCountMeta{
+ ImageGroupUnitPrice: 0.14,
+ })
+
+ require.NoError(t, err)
+ require.True(t, priceData.UsePrice)
+ require.Equal(t, 0.14, priceData.ModelPrice)
+ require.Equal(t, int(0.14*common.QuotaPerUnit), priceData.QuotaToPreConsume)
+}
+
+func TestModelPriceHelperNonImageGroupIgnoresResolutionUnitPrice(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ w := httptest.NewRecorder()
+ ctx, _ := gin.CreateTestContext(w)
+
+ info := &relaycommon.RelayInfo{
+ OriginModelName: "unknown-image-model",
+ UsingGroup: "default",
+ }
+
+ _, err := ModelPriceHelper(ctx, info, 1, &types.TokenCountMeta{
+ ImageGroupUnitPrice: 0.14,
+ })
+
+ require.Error(t, err)
+}
diff --git a/types/request_meta.go b/types/request_meta.go
index b30678bff1a7..d8c748f4f558 100644
--- a/types/request_meta.go
+++ b/types/request_meta.go
@@ -26,8 +26,9 @@ type TokenCountMeta struct {
Files []*FileMeta `json:"files,omitempty"` // List of files, each with type and content
MaxTokens int `json:"max_tokens,omitempty"` // Maximum tokens allowed in the request
- ImagePriceRatio float64 `json:"image_ratio,omitempty"` // Ratio for image size, if applicable
- ImageUnitPrice float64 `json:"image_unit_price,omitempty"`
+ ImagePriceRatio float64 `json:"image_ratio,omitempty"` // Ratio for image size, if applicable
+ ImageUnitPrice float64 `json:"image_unit_price,omitempty"`
+ ImageGroupUnitPrice float64 `json:"image_group_unit_price,omitempty"`
//IsStreaming bool `json:"is_streaming,omitempty"` // Indicates if the request is streaming
}
From 5346e64214e3e3d1911c4813e3000505bbc692de Mon Sep 17 00:00:00 2001
From: "649985538@qq.com" <649985538@qq.com>
Date: Fri, 10 Jul 2026 00:01:28 +0800
Subject: [PATCH 20/24] support gemini native image generation
---
relay/channel/gemini/adaptor.go | 129 ++++++++++++++++++++-
relay/channel/gemini/adaptor_image_test.go | 124 ++++++++++++++++++++
relay/channel/gemini/relay-gemini.go | 56 +++++++++
3 files changed, 306 insertions(+), 3 deletions(-)
create mode 100644 relay/channel/gemini/adaptor_image_test.go
diff --git a/relay/channel/gemini/adaptor.go b/relay/channel/gemini/adaptor.go
index 680c4ee484ec..b37574b14140 100644
--- a/relay/channel/gemini/adaptor.go
+++ b/relay/channel/gemini/adaptor.go
@@ -5,8 +5,10 @@ import (
"fmt"
"io"
"net/http"
+ "strconv"
"strings"
+ "github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/dto"
"github.com/QuantumNous/new-api/relay/channel"
"github.com/QuantumNous/new-api/relay/channel/openai"
@@ -59,7 +61,32 @@ func (a *Adaptor) ConvertAudioRequest(c *gin.Context, info *relaycommon.RelayInf
func (a *Adaptor) ConvertImageRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.ImageRequest) (any, error) {
if !strings.HasPrefix(info.UpstreamModelName, "imagen") {
- return nil, errors.New("not supported model for image generation, only imagen models are supported")
+ if !isGeminiNativeImageGenerationModel(info.UpstreamModelName) {
+ return nil, errors.New("not supported model for image generation, only imagen or Gemini native image models are supported")
+ }
+ if lo.FromPtrOr(request.N, uint(1)) > 1 {
+ return nil, errors.New("Gemini native image generation only supports n=1")
+ }
+ imageConfig, err := buildGeminiNativeImageConfig(request)
+ if err != nil {
+ return nil, err
+ }
+ return dto.GeminiChatRequest{
+ Contents: []dto.GeminiChatContent{
+ {
+ Role: "user",
+ Parts: []dto.GeminiPart{
+ {
+ Text: request.Prompt,
+ },
+ },
+ },
+ },
+ GenerationConfig: dto.GeminiChatGenerationConfig{
+ ResponseModalities: []string{"TEXT", "IMAGE"},
+ ImageConfig: imageConfig,
+ },
+ }, nil
}
// convert size to aspect ratio but allow user to specify aspect ratio
@@ -123,6 +150,96 @@ func (a *Adaptor) ConvertImageRequest(c *gin.Context, info *relaycommon.RelayInf
return geminiRequest, nil
}
+func isGeminiNativeImageGenerationModel(model string) bool {
+ if strings.HasPrefix(model, "imagen") {
+ return false
+ }
+ if model_setting.IsGeminiModelSupportImagine(model) {
+ return true
+ }
+ return strings.HasPrefix(model, "gemini-") &&
+ (strings.Contains(model, "-image") || strings.Contains(model, "image-generation"))
+}
+
+func buildGeminiNativeImageConfig(request dto.ImageRequest) ([]byte, error) {
+ imageSize, aspectRatio := geminiNativeImageSizeAndAspectRatio(request.Size)
+ imageConfig := map[string]interface{}{
+ "imageSize": imageSize,
+ }
+ if aspectRatio != "" {
+ imageConfig["aspectRatio"] = aspectRatio
+ }
+ imageConfigBytes, err := common.Marshal(imageConfig)
+ if err != nil {
+ return nil, fmt.Errorf("failed to marshal image config: %w", err)
+ }
+ return imageConfigBytes, nil
+}
+
+func geminiNativeImageSizeAndAspectRatio(size string) (string, string) {
+ size = strings.TrimSpace(size)
+ if size == "" || strings.EqualFold(size, "auto") {
+ return "1K", ""
+ }
+ if strings.Contains(size, ":") {
+ return "1K", size
+ }
+
+ parts := strings.Split(strings.ToLower(size), "x")
+ if len(parts) != 2 {
+ return "1K", ""
+ }
+ width, widthErr := strconv.Atoi(strings.TrimSpace(parts[0]))
+ height, heightErr := strconv.Atoi(strings.TrimSpace(parts[1]))
+ if widthErr != nil || heightErr != nil || width <= 0 || height <= 0 {
+ return "1K", ""
+ }
+
+ imageSize := "1K"
+ longEdge := width
+ if height > longEdge {
+ longEdge = height
+ }
+ if longEdge > 2048 {
+ imageSize = "4K"
+ } else if longEdge > 1024 {
+ imageSize = "2K"
+ }
+
+ return imageSize, geminiNativeAspectRatio(width, height)
+}
+
+func geminiNativeAspectRatio(width, height int) string {
+ switch fmt.Sprintf("%dx%d", width, height) {
+ case "256x256", "512x512", "1024x1024", "2048x2048", "4096x4096":
+ return "1:1"
+ case "1536x1024":
+ return "3:2"
+ case "1024x1536":
+ return "2:3"
+ case "1024x1792":
+ return "9:16"
+ case "1792x1024":
+ return "16:9"
+ }
+
+ divisor := greatestCommonDivisor(width, height)
+ return fmt.Sprintf("%d:%d", width/divisor, height/divisor)
+}
+
+func greatestCommonDivisor(a, b int) int {
+ for b != 0 {
+ a, b = b, a%b
+ }
+ if a < 0 {
+ return -a
+ }
+ if a == 0 {
+ return 1
+ }
+ return a
+}
+
func (a *Adaptor) Init(info *relaycommon.RelayInfo) {
}
@@ -259,8 +376,14 @@ func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycom
}
}
- if strings.HasPrefix(info.UpstreamModelName, "imagen") {
- return GeminiImageHandler(c, info, resp)
+ if info.RelayMode == constant.RelayModeImagesGenerations ||
+ info.RelayMode == constant.RelayModeImagesEdits {
+ if strings.HasPrefix(info.UpstreamModelName, "imagen") {
+ return GeminiImageHandler(c, info, resp)
+ }
+ if isGeminiNativeImageGenerationModel(info.UpstreamModelName) {
+ return GeminiNativeImageHandler(c, info, resp)
+ }
}
// check if the model is an embedding model
diff --git a/relay/channel/gemini/adaptor_image_test.go b/relay/channel/gemini/adaptor_image_test.go
new file mode 100644
index 000000000000..c841c951d440
--- /dev/null
+++ b/relay/channel/gemini/adaptor_image_test.go
@@ -0,0 +1,124 @@
+package gemini
+
+import (
+ "bytes"
+ "io"
+ "net/http"
+ "net/http/httptest"
+ "testing"
+
+ "github.com/QuantumNous/new-api/common"
+ "github.com/QuantumNous/new-api/dto"
+ relaycommon "github.com/QuantumNous/new-api/relay/common"
+ "github.com/gin-gonic/gin"
+ "github.com/stretchr/testify/require"
+)
+
+func TestConvertImageRequestGeminiNativeImageModel(t *testing.T) {
+ t.Parallel()
+
+ gin.SetMode(gin.TestMode)
+ c, _ := gin.CreateTestContext(httptest.NewRecorder())
+
+ n := uint(1)
+ converted, err := (&Adaptor{}).ConvertImageRequest(c, &relaycommon.RelayInfo{
+ ChannelMeta: &relaycommon.ChannelMeta{
+ UpstreamModelName: "gemini-3.1-flash-image",
+ },
+ }, dto.ImageRequest{
+ Prompt: "draw a small red house",
+ Size: "2048x2048",
+ N: &n,
+ })
+
+ require.NoError(t, err)
+ geminiRequest, ok := converted.(dto.GeminiChatRequest)
+ require.True(t, ok)
+ require.Len(t, geminiRequest.Contents, 1)
+ require.Equal(t, "user", geminiRequest.Contents[0].Role)
+ require.Equal(t, "draw a small red house", geminiRequest.Contents[0].Parts[0].Text)
+ require.Equal(t, []string{"TEXT", "IMAGE"}, geminiRequest.GenerationConfig.ResponseModalities)
+
+ var imageConfig map[string]string
+ require.NoError(t, common.Unmarshal(geminiRequest.GenerationConfig.ImageConfig, &imageConfig))
+ require.Equal(t, "2K", imageConfig["imageSize"])
+ require.Equal(t, "1:1", imageConfig["aspectRatio"])
+}
+
+func TestConvertImageRequestGeminiNativeImageRejectsMultipleImages(t *testing.T) {
+ t.Parallel()
+
+ gin.SetMode(gin.TestMode)
+ c, _ := gin.CreateTestContext(httptest.NewRecorder())
+
+ n := uint(2)
+ _, err := (&Adaptor{}).ConvertImageRequest(c, &relaycommon.RelayInfo{
+ ChannelMeta: &relaycommon.ChannelMeta{
+ UpstreamModelName: "gemini-3.1-flash-image",
+ },
+ }, dto.ImageRequest{
+ Prompt: "draw a small red house",
+ Size: "1024x1024",
+ N: &n,
+ })
+
+ require.ErrorContains(t, err, "only supports n=1")
+}
+
+func TestGeminiNativeImageHandlerConvertsInlineImageToOpenAIImageResponse(t *testing.T) {
+ t.Parallel()
+
+ gin.SetMode(gin.TestMode)
+ recorder := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(recorder)
+ c.Request = httptest.NewRequest(http.MethodPost, "/v1/images/generations", nil)
+
+ payload := dto.GeminiChatResponse{
+ Candidates: []dto.GeminiChatCandidate{
+ {
+ Content: dto.GeminiChatContent{
+ Role: "model",
+ Parts: []dto.GeminiPart{
+ {Text: "revised prompt"},
+ {
+ InlineData: &dto.GeminiInlineData{
+ MimeType: "image/png",
+ Data: "aW1hZ2UtYnl0ZXM=",
+ },
+ },
+ },
+ },
+ },
+ },
+ UsageMetadata: dto.GeminiUsageMetadata{
+ PromptTokenCount: 11,
+ CandidatesTokenCount: 22,
+ TotalTokenCount: 33,
+ },
+ }
+ body, err := common.Marshal(payload)
+ require.NoError(t, err)
+
+ info := &relaycommon.RelayInfo{
+ ChannelMeta: &relaycommon.ChannelMeta{
+ UpstreamModelName: "gemini-3.1-flash-image",
+ },
+ }
+ usage, newAPIError := GeminiNativeImageHandler(c, info, &http.Response{
+ StatusCode: http.StatusOK,
+ Body: io.NopCloser(bytes.NewReader(body)),
+ })
+
+ require.Nil(t, newAPIError)
+ require.NotNil(t, usage)
+ require.Equal(t, 11, usage.PromptTokens)
+ require.Equal(t, 22, usage.CompletionTokens)
+ require.Equal(t, 33, usage.TotalTokens)
+ require.Equal(t, float64(1), info.PriceData.OtherRatios["n"])
+
+ var openAIResponse dto.ImageResponse
+ require.NoError(t, common.Unmarshal(recorder.Body.Bytes(), &openAIResponse))
+ require.Len(t, openAIResponse.Data, 1)
+ require.Equal(t, "aW1hZ2UtYnl0ZXM=", openAIResponse.Data[0].B64Json)
+ require.Equal(t, "revised prompt", openAIResponse.Data[0].RevisedPrompt)
+}
diff --git a/relay/channel/gemini/relay-gemini.go b/relay/channel/gemini/relay-gemini.go
index 69175e76efc6..bf0b188f7ff0 100644
--- a/relay/channel/gemini/relay-gemini.go
+++ b/relay/channel/gemini/relay-gemini.go
@@ -1582,6 +1582,62 @@ func GeminiImageHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.
return usage, nil
}
+func GeminiNativeImageHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
+ responseBody, readErr := io.ReadAll(resp.Body)
+ if readErr != nil {
+ return nil, types.NewOpenAIError(readErr, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
+ }
+ service.CloseResponseBodyGracefully(resp)
+
+ var geminiResponse dto.GeminiChatResponse
+ if jsonErr := common.Unmarshal(responseBody, &geminiResponse); jsonErr != nil {
+ return nil, types.NewOpenAIError(jsonErr, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
+ }
+
+ openAIResponse := dto.ImageResponse{
+ Created: common.GetTimestamp(),
+ }
+ var revisedPrompts []string
+ for _, candidate := range geminiResponse.Candidates {
+ for _, part := range candidate.Content.Parts {
+ if strings.TrimSpace(part.Text) != "" {
+ revisedPrompts = append(revisedPrompts, strings.TrimSpace(part.Text))
+ }
+ if part.InlineData == nil ||
+ part.InlineData.Data == "" ||
+ !strings.HasPrefix(strings.ToLower(part.InlineData.MimeType), "image/") {
+ continue
+ }
+ openAIResponse.Data = append(openAIResponse.Data, dto.ImageData{
+ B64Json: part.InlineData.Data,
+ })
+ }
+ }
+
+ if len(openAIResponse.Data) == 0 {
+ return nil, types.NewOpenAIError(errors.New("no images generated"), types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
+ }
+ if len(revisedPrompts) > 0 {
+ revisedPrompt := strings.Join(revisedPrompts, "\n")
+ for i := range openAIResponse.Data {
+ openAIResponse.Data[i].RevisedPrompt = revisedPrompt
+ }
+ }
+ info.PriceData.AddOtherRatio("n", float64(len(openAIResponse.Data)))
+
+ jsonResponse, jsonErr := common.Marshal(openAIResponse)
+ if jsonErr != nil {
+ return nil, types.NewError(jsonErr, types.ErrorCodeBadResponseBody)
+ }
+
+ c.Writer.Header().Set("Content-Type", "application/json")
+ c.Writer.WriteHeader(resp.StatusCode)
+ _, _ = c.Writer.Write(jsonResponse)
+
+ usage := buildUsageFromGeminiMetadata(geminiResponse.UsageMetadata, info.GetEstimatePromptTokens())
+ return &usage, nil
+}
+
type GeminiModelsResponse struct {
Models []dto.GeminiModel `json:"models"`
NextPageToken string `json:"nextPageToken"`
From 8012ee71b1d37abfccfc8542eb75a5c0f271a65a Mon Sep 17 00:00:00 2001
From: "649985538@qq.com" <649985538@qq.com>
Date: Fri, 10 Jul 2026 00:27:08 +0800
Subject: [PATCH 21/24] show image generations in drawing logs
---
controller/midjourney.go | 16 +-
model/midjourney.go | 139 ++++++++++++++++++
model/midjourney_test.go | 99 +++++++++++++
model/task_cas_test.go | 1 +
.../table/mj-logs/MjLogsColumnDefs.jsx | 6 +
5 files changed, 255 insertions(+), 6 deletions(-)
create mode 100644 model/midjourney_test.go
diff --git a/controller/midjourney.go b/controller/midjourney.go
index 69aa5ccd431f..22a3c8fd65f5 100644
--- a/controller/midjourney.go
+++ b/controller/midjourney.go
@@ -265,12 +265,14 @@ func GetAllMidjourney(c *gin.Context) {
EndTimestamp: c.Query("end_timestamp"),
}
- items := model.GetAllTasks(pageInfo.GetStartIdx(), pageInfo.GetPageSize(), queryParams)
- total := model.CountAllTasks(queryParams)
+ items := model.GetAllDrawingLogs(pageInfo.GetStartIdx(), pageInfo.GetPageSize(), queryParams)
+ total := model.CountAllDrawingLogs(queryParams)
if setting.MjForwardUrlEnabled {
for i, midjourney := range items {
- midjourney.ImageUrl = system_setting.ServerAddress + "/mj/image/" + midjourney.MjId
+ if midjourney.Id > 0 {
+ midjourney.ImageUrl = system_setting.ServerAddress + "/mj/image/" + midjourney.MjId
+ }
items[i] = midjourney
}
}
@@ -290,12 +292,14 @@ func GetUserMidjourney(c *gin.Context) {
EndTimestamp: c.Query("end_timestamp"),
}
- items := model.GetAllUserTask(userId, pageInfo.GetStartIdx(), pageInfo.GetPageSize(), queryParams)
- total := model.CountAllUserTask(userId, queryParams)
+ items := model.GetAllUserDrawingLogs(userId, pageInfo.GetStartIdx(), pageInfo.GetPageSize(), queryParams)
+ total := model.CountAllUserDrawingLogs(userId, queryParams)
if setting.MjForwardUrlEnabled {
for i, midjourney := range items {
- midjourney.ImageUrl = system_setting.ServerAddress + "/mj/image/" + midjourney.MjId
+ if midjourney.Id > 0 {
+ midjourney.ImageUrl = system_setting.ServerAddress + "/mj/image/" + midjourney.MjId
+ }
items[i] = midjourney
}
}
diff --git a/model/midjourney.go b/model/midjourney.go
index e1a8d772b068..3d2bc7038d38 100644
--- a/model/midjourney.go
+++ b/model/midjourney.go
@@ -1,5 +1,13 @@
package model
+import (
+ "fmt"
+ "sort"
+ "strconv"
+
+ "gorm.io/gorm"
+)
+
type Midjourney struct {
Id int `json:"id"`
Code int `json:"code"`
@@ -218,3 +226,134 @@ func CountAllUserTask(userId int, queryParams TaskQueryParams) int64 {
_ = query.Count(&total).Error
return total
}
+
+func GetAllDrawingLogs(startIdx int, num int, queryParams TaskQueryParams) []*Midjourney {
+ limit := startIdx + num
+ items := append(
+ GetAllTasks(0, limit, queryParams),
+ GetAllImageGenerationLogTasks(0, limit, queryParams, nil)...,
+ )
+ return paginateDrawingLogs(items, startIdx, num)
+}
+
+func GetAllUserDrawingLogs(userId int, startIdx int, num int, queryParams TaskQueryParams) []*Midjourney {
+ limit := startIdx + num
+ items := append(
+ GetAllUserTask(userId, 0, limit, queryParams),
+ GetAllImageGenerationLogTasks(0, limit, queryParams, &userId)...,
+ )
+ return paginateDrawingLogs(items, startIdx, num)
+}
+
+func CountAllDrawingLogs(queryParams TaskQueryParams) int64 {
+ return CountAllTasks(queryParams) + CountAllImageGenerationLogTasks(queryParams, nil)
+}
+
+func CountAllUserDrawingLogs(userId int, queryParams TaskQueryParams) int64 {
+ return CountAllUserTask(userId, queryParams) + CountAllImageGenerationLogTasks(queryParams, &userId)
+}
+
+func GetAllImageGenerationLogTasks(startIdx int, num int, queryParams TaskQueryParams, userId *int) []*Midjourney {
+ var logs []*Log
+ tx := imageGenerationLogTaskQuery(queryParams, userId)
+ err := tx.Order("logs.created_at desc, logs.id desc").Limit(num).Offset(startIdx).Find(&logs).Error
+ if err != nil {
+ return nil
+ }
+
+ items := make([]*Midjourney, 0, len(logs))
+ for _, log := range logs {
+ items = append(items, imageGenerationLogToMidjourney(log))
+ }
+ return items
+}
+
+func CountAllImageGenerationLogTasks(queryParams TaskQueryParams, userId *int) int64 {
+ var total int64
+ _ = imageGenerationLogTaskQuery(queryParams, userId).Count(&total).Error
+ return total
+}
+
+func imageGenerationLogTaskQuery(queryParams TaskQueryParams, userId *int) *gorm.DB {
+ tx := LOG_DB.Model(&Log{}).Where("logs.type = ?", LogTypeConsume).
+ Where("logs.other LIKE ?", "%\"request_path\":\"/v1/images/generations\"%")
+
+ if userId != nil {
+ tx = tx.Where("logs.user_id = ?", *userId)
+ }
+ if queryParams.ChannelID != "" {
+ tx = tx.Where("logs.channel_id = ?", queryParams.ChannelID)
+ }
+ if queryParams.MjID != "" {
+ tx = tx.Where("logs.request_id = ?", queryParams.MjID)
+ }
+ if startTimestamp := taskTimestampMillisToSeconds(queryParams.StartTimestamp); startTimestamp > 0 {
+ tx = tx.Where("logs.created_at >= ?", startTimestamp)
+ }
+ if endTimestamp := taskTimestampMillisToSeconds(queryParams.EndTimestamp); endTimestamp > 0 {
+ tx = tx.Where("logs.created_at <= ?", endTimestamp)
+ }
+
+ return tx
+}
+
+func imageGenerationLogToMidjourney(log *Log) *Midjourney {
+ mjId := log.RequestId
+ if mjId == "" {
+ mjId = fmt.Sprintf("image_log_%d", log.Id)
+ }
+
+ finishTime := log.CreatedAt * 1000
+ if log.UseTime > 0 {
+ finishTime = (log.CreatedAt + int64(log.UseTime)) * 1000
+ }
+
+ return &Midjourney{
+ Id: -log.Id,
+ Code: 1,
+ UserId: log.UserId,
+ Action: "IMAGE_GENERATION",
+ MjId: mjId,
+ Prompt: log.Content,
+ PromptEn: log.ModelName,
+ SubmitTime: log.CreatedAt * 1000,
+ StartTime: log.CreatedAt * 1000,
+ FinishTime: finishTime,
+ Status: "SUCCESS",
+ Progress: "100%",
+ ChannelId: log.ChannelId,
+ Quota: log.Quota,
+ }
+}
+
+func paginateDrawingLogs(items []*Midjourney, startIdx int, num int) []*Midjourney {
+ sort.SliceStable(items, func(i, j int) bool {
+ if items[i].SubmitTime == items[j].SubmitTime {
+ return items[i].Id > items[j].Id
+ }
+ return items[i].SubmitTime > items[j].SubmitTime
+ })
+
+ if startIdx >= len(items) {
+ return []*Midjourney{}
+ }
+ endIdx := startIdx + num
+ if endIdx > len(items) {
+ endIdx = len(items)
+ }
+ return items[startIdx:endIdx]
+}
+
+func taskTimestampMillisToSeconds(raw string) int64 {
+ if raw == "" {
+ return 0
+ }
+ timestamp, err := strconv.ParseInt(raw, 10, 64)
+ if err != nil || timestamp <= 0 {
+ return 0
+ }
+ if timestamp > 1_000_000_000_000 {
+ return timestamp / 1000
+ }
+ return timestamp
+}
diff --git a/model/midjourney_test.go b/model/midjourney_test.go
new file mode 100644
index 000000000000..c490f20fdea5
--- /dev/null
+++ b/model/midjourney_test.go
@@ -0,0 +1,99 @@
+package model
+
+import (
+ "testing"
+
+ "github.com/QuantumNous/new-api/common"
+ "github.com/stretchr/testify/require"
+)
+
+func TestGetAllUserDrawingLogsIncludesImageGenerationLogs(t *testing.T) {
+ truncateTables(t)
+
+ require.NoError(t, DB.Create(&Midjourney{
+ Id: 1,
+ UserId: 1,
+ Action: "IMAGINE",
+ MjId: "mj_old",
+ Prompt: "old mj prompt",
+ SubmitTime: 1000,
+ Status: "SUCCESS",
+ Progress: "100%",
+ ChannelId: 9,
+ }).Error)
+
+ require.NoError(t, LOG_DB.Create(&Log{
+ Id: 10,
+ UserId: 1,
+ CreatedAt: 2,
+ Type: LogTypeConsume,
+ Content: "大小 1024x1024, 品质 standard, 生成数量 1",
+ ModelName: "gemini-3.1-flash-image",
+ Quota: 50000,
+ UseTime: 3,
+ ChannelId: 23,
+ RequestId: "req_image",
+ Other: common.MapToJsonStr(map[string]interface{}{
+ "request_path": "/v1/images/generations",
+ "model_price": 0.1,
+ }),
+ }).Error)
+
+ require.NoError(t, LOG_DB.Create(&Log{
+ Id: 11,
+ UserId: 1,
+ CreatedAt: 3,
+ Type: LogTypeConsume,
+ Content: "chat",
+ ModelName: "gpt-4o",
+ Other: common.MapToJsonStr(map[string]interface{}{
+ "request_path": "/v1/chat/completions",
+ }),
+ }).Error)
+
+ items := GetAllUserDrawingLogs(1, 0, 10, TaskQueryParams{})
+ require.Len(t, items, 2)
+ require.Equal(t, "req_image", items[0].MjId)
+ require.Equal(t, "IMAGE_GENERATION", items[0].Action)
+ require.Equal(t, "SUCCESS", items[0].Status)
+ require.Equal(t, "100%", items[0].Progress)
+ require.Equal(t, int64(2000), items[0].SubmitTime)
+ require.Equal(t, int64(5000), items[0].FinishTime)
+ require.Equal(t, "gemini-3.1-flash-image", items[0].PromptEn)
+ require.Equal(t, 50000, items[0].Quota)
+ require.Equal(t, "mj_old", items[1].MjId)
+ require.Equal(t, int64(2), CountAllUserDrawingLogs(1, TaskQueryParams{}))
+}
+
+func TestGetAllUserDrawingLogsFiltersImageGenerationByRequestID(t *testing.T) {
+ truncateTables(t)
+
+ require.NoError(t, LOG_DB.Create(&Log{
+ Id: 20,
+ UserId: 1,
+ CreatedAt: 2,
+ Type: LogTypeConsume,
+ Content: "大小 2048x2048, 品质 standard, 生成数量 1",
+ ModelName: "gpt-image-2",
+ RequestId: "req_match",
+ Other: common.MapToJsonStr(map[string]interface{}{
+ "request_path": "/v1/images/generations",
+ }),
+ }).Error)
+ require.NoError(t, LOG_DB.Create(&Log{
+ Id: 21,
+ UserId: 1,
+ CreatedAt: 3,
+ Type: LogTypeConsume,
+ Content: "大小 4096x4096, 品质 standard, 生成数量 1",
+ ModelName: "gpt-image-2",
+ RequestId: "req_other",
+ Other: common.MapToJsonStr(map[string]interface{}{
+ "request_path": "/v1/images/generations",
+ }),
+ }).Error)
+
+ items := GetAllUserDrawingLogs(1, 0, 10, TaskQueryParams{MjID: "req_match"})
+ require.Len(t, items, 1)
+ require.Equal(t, "req_match", items[0].MjId)
+}
diff --git a/model/task_cas_test.go b/model/task_cas_test.go
index ba34a73291bc..189d07c64207 100644
--- a/model/task_cas_test.go
+++ b/model/task_cas_test.go
@@ -35,6 +35,7 @@ func TestMain(m *testing.M) {
if err := db.AutoMigrate(
&Task{},
+ &Midjourney{},
&User{},
&Token{},
&Log{},
diff --git a/web/src/components/table/mj-logs/MjLogsColumnDefs.jsx b/web/src/components/table/mj-logs/MjLogsColumnDefs.jsx
index 9fa26efe0b90..c9ce662a3a9d 100644
--- a/web/src/components/table/mj-logs/MjLogsColumnDefs.jsx
+++ b/web/src/components/table/mj-logs/MjLogsColumnDefs.jsx
@@ -74,6 +74,12 @@ function renderType(type, t) {
{t('绘图')}
);
+ case 'IMAGE_GENERATION':
+ return (
+
}>
+ {t('图片生成')}
+
+ );
case 'UPSCALE':
return (
}>
From 533f49c7bfe2719fd8108e91b1499311887bc0f6 Mon Sep 17 00:00:00 2001
From: "649985538@qq.com" <649985538@qq.com>
Date: Fri, 10 Jul 2026 00:55:50 +0800
Subject: [PATCH 22/24] store image generation results locally
---
controller/image_generation.go | 54 +++++
main.go | 3 +
model/image_generation.go | 145 ++++++++++++++
model/main.go | 2 +
model/midjourney.go | 84 +-------
model/midjourney_test.go | 56 ++++--
model/task_cas_test.go | 3 +
relay/image_handler.go | 26 ++-
router/api-router.go | 3 +
service/image_generation_storage.go | 184 ++++++++++++++++++
service/image_generation_storage_test.go | 104 ++++++++++
service/task_billing_test.go | 1 +
service/text_quota.go | 3 +-
.../table/mj-logs/MjLogsColumnDefs.jsx | 6 +
14 files changed, 571 insertions(+), 103 deletions(-)
create mode 100644 controller/image_generation.go
create mode 100644 model/image_generation.go
create mode 100644 service/image_generation_storage.go
create mode 100644 service/image_generation_storage_test.go
diff --git a/controller/image_generation.go b/controller/image_generation.go
new file mode 100644
index 000000000000..cc5f7118082b
--- /dev/null
+++ b/controller/image_generation.go
@@ -0,0 +1,54 @@
+package controller
+
+import (
+ "net/http"
+ "os"
+ "strconv"
+
+ "github.com/QuantumNous/new-api/common"
+ "github.com/QuantumNous/new-api/model"
+ "github.com/QuantumNous/new-api/service"
+
+ "github.com/gin-gonic/gin"
+)
+
+func GetImageGenerationContent(c *gin.Context) {
+ id, err := strconv.Atoi(c.Param("id"))
+ if err != nil || id <= 0 {
+ c.Status(http.StatusNotFound)
+ return
+ }
+
+ record, err := model.GetImageGenerationByID(id)
+ if err != nil || record == nil {
+ c.Status(http.StatusNotFound)
+ return
+ }
+
+ role := c.GetInt("role")
+ userID := c.GetInt("id")
+ if role < common.RoleAdminUser && record.UserId != userID {
+ c.Status(http.StatusForbidden)
+ return
+ }
+ if record.Status != model.ImageGenerationStatusSuccess || record.FilePath == "" {
+ c.Status(http.StatusGone)
+ return
+ }
+
+ absolutePath := service.GetImageGenerationAbsolutePath(record)
+ if absolutePath == "" {
+ c.Status(http.StatusGone)
+ return
+ }
+ if _, err := os.Stat(absolutePath); err != nil {
+ c.Status(http.StatusGone)
+ return
+ }
+
+ if record.MimeType != "" {
+ c.Header("Content-Type", record.MimeType)
+ }
+ c.Header("Cache-Control", "private, max-age=3600")
+ c.File(absolutePath)
+}
diff --git a/main.go b/main.go
index d8587e70e6bd..f162122863cb 100644
--- a/main.go
+++ b/main.go
@@ -112,6 +112,9 @@ func main() {
// Subscription quota reset task (daily/weekly/monthly/custom)
service.StartSubscriptionQuotaResetTask()
+ // Local image generation result retention cleanup.
+ service.StartImageGenerationCleanupTask()
+
// Optional AistarsLab video model/price sync task.
service.StartAistarsLabConfigSyncTask()
diff --git a/model/image_generation.go b/model/image_generation.go
new file mode 100644
index 000000000000..892c57e24aa5
--- /dev/null
+++ b/model/image_generation.go
@@ -0,0 +1,145 @@
+package model
+
+import (
+ "strconv"
+
+ "gorm.io/gorm"
+)
+
+const (
+ ImageGenerationStatusSuccess = "SUCCESS"
+ ImageGenerationStatusExpired = "EXPIRED"
+)
+
+type ImageGeneration struct {
+ Id int `json:"id"`
+ UserId int `json:"user_id" gorm:"index"`
+ TokenId int `json:"token_id" gorm:"index"`
+ ChannelId int `json:"channel_id" gorm:"index"`
+ RequestId string `json:"request_id" gorm:"type:varchar(64);index"`
+ ImageIndex int `json:"image_index" gorm:"index"`
+ ModelName string `json:"model_name" gorm:"index"`
+ Prompt string `json:"prompt" gorm:"type:text"`
+ Size string `json:"size" gorm:"type:varchar(64)"`
+ Quality string `json:"quality" gorm:"type:varchar(64)"`
+ Quota int `json:"quota"`
+ FilePath string `json:"file_path" gorm:"type:text"`
+ MimeType string `json:"mime_type" gorm:"type:varchar(64)"`
+ Status string `json:"status" gorm:"type:varchar(20);index"`
+ Group string `json:"group" gorm:"index"`
+ CreatedAt int64 `json:"created_at" gorm:"bigint;index"`
+ ExpireAt int64 `json:"expire_at" gorm:"bigint;index"`
+}
+
+func InsertImageGeneration(record *ImageGeneration) error {
+ return DB.Create(record).Error
+}
+
+func GetImageGenerationByID(id int) (*ImageGeneration, error) {
+ var record ImageGeneration
+ err := DB.Where("id = ?", id).First(&record).Error
+ if err != nil {
+ return nil, err
+ }
+ return &record, nil
+}
+
+func GetExpiredImageGenerations(now int64, limit int) ([]*ImageGeneration, error) {
+ var records []*ImageGeneration
+ err := DB.Where("status = ? AND expire_at <= ?", ImageGenerationStatusSuccess, now).
+ Limit(limit).
+ Find(&records).Error
+ return records, err
+}
+
+func MarkImageGenerationExpired(id int) error {
+ return DB.Model(&ImageGeneration{}).
+ Where("id = ?", id).
+ Updates(map[string]interface{}{
+ "status": ImageGenerationStatusExpired,
+ "file_path": "",
+ }).Error
+}
+
+func imageGenerationQuery(queryParams TaskQueryParams, userId *int) *gorm.DB {
+ tx := DB.Model(&ImageGeneration{})
+ if userId != nil {
+ tx = tx.Where("user_id = ?", *userId)
+ }
+ if queryParams.ChannelID != "" {
+ tx = tx.Where("channel_id = ?", queryParams.ChannelID)
+ }
+ if queryParams.MjID != "" {
+ tx = tx.Where("request_id = ?", queryParams.MjID)
+ }
+ if startTimestamp := taskTimestampMillisToSeconds(queryParams.StartTimestamp); startTimestamp > 0 {
+ tx = tx.Where("created_at >= ?", startTimestamp)
+ }
+ if endTimestamp := taskTimestampMillisToSeconds(queryParams.EndTimestamp); endTimestamp > 0 {
+ tx = tx.Where("created_at <= ?", endTimestamp)
+ }
+ return tx
+}
+
+func GetAllImageGenerationTasks(startIdx int, num int, queryParams TaskQueryParams, userId *int) []*Midjourney {
+ var records []*ImageGeneration
+ err := imageGenerationQuery(queryParams, userId).
+ Order("created_at desc, id desc").
+ Limit(num).
+ Offset(startIdx).
+ Find(&records).Error
+ if err != nil {
+ return nil
+ }
+
+ items := make([]*Midjourney, 0, len(records))
+ for _, record := range records {
+ items = append(items, imageGenerationToMidjourney(record))
+ }
+ return items
+}
+
+func CountAllImageGenerationTasks(queryParams TaskQueryParams, userId *int) int64 {
+ var total int64
+ _ = imageGenerationQuery(queryParams, userId).Count(&total).Error
+ return total
+}
+
+func imageGenerationToMidjourney(record *ImageGeneration) *Midjourney {
+ imageURL := ""
+ failReason := ""
+ status := record.Status
+ if status == "" {
+ status = ImageGenerationStatusSuccess
+ }
+ if status == ImageGenerationStatusSuccess && record.FilePath != "" {
+ imageURL = "/api/image-generations/" + strconv.Itoa(record.Id) + "/content"
+ }
+ if status == ImageGenerationStatusExpired {
+ failReason = "图片已过期"
+ }
+
+ mjID := record.RequestId
+ if record.ImageIndex > 0 {
+ mjID = mjID + "#" + strconv.Itoa(record.ImageIndex+1)
+ }
+
+ return &Midjourney{
+ Id: -record.Id,
+ Code: 1,
+ UserId: record.UserId,
+ Action: "IMAGE_GENERATION",
+ MjId: mjID,
+ Prompt: record.Prompt,
+ PromptEn: record.ModelName,
+ SubmitTime: record.CreatedAt * 1000,
+ StartTime: record.CreatedAt * 1000,
+ FinishTime: record.CreatedAt * 1000,
+ ImageUrl: imageURL,
+ Status: status,
+ Progress: "100%",
+ FailReason: failReason,
+ ChannelId: record.ChannelId,
+ Quota: record.Quota,
+ }
+}
diff --git a/model/main.go b/model/main.go
index f37cb667cd43..291886930999 100644
--- a/model/main.go
+++ b/model/main.go
@@ -265,6 +265,7 @@ func migrateDB() error {
&Ability{},
&Log{},
&Midjourney{},
+ &ImageGeneration{},
&TopUp{},
&QuotaData{},
&Task{},
@@ -313,6 +314,7 @@ func migrateDBFast() error {
{&Ability{}, "Ability"},
{&Log{}, "Log"},
{&Midjourney{}, "Midjourney"},
+ {&ImageGeneration{}, "ImageGeneration"},
{&TopUp{}, "TopUp"},
{&QuotaData{}, "QuotaData"},
{&Task{}, "Task"},
diff --git a/model/midjourney.go b/model/midjourney.go
index 3d2bc7038d38..ed63150646a9 100644
--- a/model/midjourney.go
+++ b/model/midjourney.go
@@ -1,11 +1,8 @@
package model
import (
- "fmt"
"sort"
"strconv"
-
- "gorm.io/gorm"
)
type Midjourney struct {
@@ -231,7 +228,7 @@ func GetAllDrawingLogs(startIdx int, num int, queryParams TaskQueryParams) []*Mi
limit := startIdx + num
items := append(
GetAllTasks(0, limit, queryParams),
- GetAllImageGenerationLogTasks(0, limit, queryParams, nil)...,
+ GetAllImageGenerationTasks(0, limit, queryParams, nil)...,
)
return paginateDrawingLogs(items, startIdx, num)
}
@@ -240,90 +237,17 @@ func GetAllUserDrawingLogs(userId int, startIdx int, num int, queryParams TaskQu
limit := startIdx + num
items := append(
GetAllUserTask(userId, 0, limit, queryParams),
- GetAllImageGenerationLogTasks(0, limit, queryParams, &userId)...,
+ GetAllImageGenerationTasks(0, limit, queryParams, &userId)...,
)
return paginateDrawingLogs(items, startIdx, num)
}
func CountAllDrawingLogs(queryParams TaskQueryParams) int64 {
- return CountAllTasks(queryParams) + CountAllImageGenerationLogTasks(queryParams, nil)
+ return CountAllTasks(queryParams) + CountAllImageGenerationTasks(queryParams, nil)
}
func CountAllUserDrawingLogs(userId int, queryParams TaskQueryParams) int64 {
- return CountAllUserTask(userId, queryParams) + CountAllImageGenerationLogTasks(queryParams, &userId)
-}
-
-func GetAllImageGenerationLogTasks(startIdx int, num int, queryParams TaskQueryParams, userId *int) []*Midjourney {
- var logs []*Log
- tx := imageGenerationLogTaskQuery(queryParams, userId)
- err := tx.Order("logs.created_at desc, logs.id desc").Limit(num).Offset(startIdx).Find(&logs).Error
- if err != nil {
- return nil
- }
-
- items := make([]*Midjourney, 0, len(logs))
- for _, log := range logs {
- items = append(items, imageGenerationLogToMidjourney(log))
- }
- return items
-}
-
-func CountAllImageGenerationLogTasks(queryParams TaskQueryParams, userId *int) int64 {
- var total int64
- _ = imageGenerationLogTaskQuery(queryParams, userId).Count(&total).Error
- return total
-}
-
-func imageGenerationLogTaskQuery(queryParams TaskQueryParams, userId *int) *gorm.DB {
- tx := LOG_DB.Model(&Log{}).Where("logs.type = ?", LogTypeConsume).
- Where("logs.other LIKE ?", "%\"request_path\":\"/v1/images/generations\"%")
-
- if userId != nil {
- tx = tx.Where("logs.user_id = ?", *userId)
- }
- if queryParams.ChannelID != "" {
- tx = tx.Where("logs.channel_id = ?", queryParams.ChannelID)
- }
- if queryParams.MjID != "" {
- tx = tx.Where("logs.request_id = ?", queryParams.MjID)
- }
- if startTimestamp := taskTimestampMillisToSeconds(queryParams.StartTimestamp); startTimestamp > 0 {
- tx = tx.Where("logs.created_at >= ?", startTimestamp)
- }
- if endTimestamp := taskTimestampMillisToSeconds(queryParams.EndTimestamp); endTimestamp > 0 {
- tx = tx.Where("logs.created_at <= ?", endTimestamp)
- }
-
- return tx
-}
-
-func imageGenerationLogToMidjourney(log *Log) *Midjourney {
- mjId := log.RequestId
- if mjId == "" {
- mjId = fmt.Sprintf("image_log_%d", log.Id)
- }
-
- finishTime := log.CreatedAt * 1000
- if log.UseTime > 0 {
- finishTime = (log.CreatedAt + int64(log.UseTime)) * 1000
- }
-
- return &Midjourney{
- Id: -log.Id,
- Code: 1,
- UserId: log.UserId,
- Action: "IMAGE_GENERATION",
- MjId: mjId,
- Prompt: log.Content,
- PromptEn: log.ModelName,
- SubmitTime: log.CreatedAt * 1000,
- StartTime: log.CreatedAt * 1000,
- FinishTime: finishTime,
- Status: "SUCCESS",
- Progress: "100%",
- ChannelId: log.ChannelId,
- Quota: log.Quota,
- }
+ return CountAllUserTask(userId, queryParams) + CountAllImageGenerationTasks(queryParams, &userId)
}
func paginateDrawingLogs(items []*Midjourney, startIdx int, num int) []*Midjourney {
diff --git a/model/midjourney_test.go b/model/midjourney_test.go
index c490f20fdea5..3771f3be6c3b 100644
--- a/model/midjourney_test.go
+++ b/model/midjourney_test.go
@@ -4,6 +4,7 @@ import (
"testing"
"github.com/QuantumNous/new-api/common"
+
"github.com/stretchr/testify/require"
)
@@ -22,21 +23,17 @@ func TestGetAllUserDrawingLogsIncludesImageGenerationLogs(t *testing.T) {
ChannelId: 9,
}).Error)
- require.NoError(t, LOG_DB.Create(&Log{
+ require.NoError(t, DB.Create(&ImageGeneration{
Id: 10,
UserId: 1,
CreatedAt: 2,
- Type: LogTypeConsume,
- Content: "大小 1024x1024, 品质 standard, 生成数量 1",
+ Status: ImageGenerationStatusSuccess,
+ Prompt: "a red cube",
ModelName: "gemini-3.1-flash-image",
Quota: 50000,
- UseTime: 3,
ChannelId: 23,
RequestId: "req_image",
- Other: common.MapToJsonStr(map[string]interface{}{
- "request_path": "/v1/images/generations",
- "model_price": 0.1,
- }),
+ FilePath: "20260710/user-1/req_image-0.png",
}).Error)
require.NoError(t, LOG_DB.Create(&Log{
@@ -58,7 +55,9 @@ func TestGetAllUserDrawingLogsIncludesImageGenerationLogs(t *testing.T) {
require.Equal(t, "SUCCESS", items[0].Status)
require.Equal(t, "100%", items[0].Progress)
require.Equal(t, int64(2000), items[0].SubmitTime)
- require.Equal(t, int64(5000), items[0].FinishTime)
+ require.Equal(t, int64(2000), items[0].FinishTime)
+ require.Equal(t, "/api/image-generations/10/content", items[0].ImageUrl)
+ require.Equal(t, "a red cube", items[0].Prompt)
require.Equal(t, "gemini-3.1-flash-image", items[0].PromptEn)
require.Equal(t, 50000, items[0].Quota)
require.Equal(t, "mj_old", items[1].MjId)
@@ -68,32 +67,47 @@ func TestGetAllUserDrawingLogsIncludesImageGenerationLogs(t *testing.T) {
func TestGetAllUserDrawingLogsFiltersImageGenerationByRequestID(t *testing.T) {
truncateTables(t)
- require.NoError(t, LOG_DB.Create(&Log{
+ require.NoError(t, DB.Create(&ImageGeneration{
Id: 20,
UserId: 1,
CreatedAt: 2,
- Type: LogTypeConsume,
- Content: "大小 2048x2048, 品质 standard, 生成数量 1",
+ Status: ImageGenerationStatusSuccess,
+ Prompt: "match",
ModelName: "gpt-image-2",
RequestId: "req_match",
- Other: common.MapToJsonStr(map[string]interface{}{
- "request_path": "/v1/images/generations",
- }),
}).Error)
- require.NoError(t, LOG_DB.Create(&Log{
+ require.NoError(t, DB.Create(&ImageGeneration{
Id: 21,
UserId: 1,
CreatedAt: 3,
- Type: LogTypeConsume,
- Content: "大小 4096x4096, 品质 standard, 生成数量 1",
+ Status: ImageGenerationStatusSuccess,
+ Prompt: "other",
ModelName: "gpt-image-2",
RequestId: "req_other",
- Other: common.MapToJsonStr(map[string]interface{}{
- "request_path": "/v1/images/generations",
- }),
}).Error)
items := GetAllUserDrawingLogs(1, 0, 10, TaskQueryParams{MjID: "req_match"})
require.Len(t, items, 1)
require.Equal(t, "req_match", items[0].MjId)
}
+
+func TestGetAllUserDrawingLogsShowsExpiredImageGeneration(t *testing.T) {
+ truncateTables(t)
+
+ require.NoError(t, DB.Create(&ImageGeneration{
+ Id: 30,
+ UserId: 1,
+ CreatedAt: 2,
+ Status: ImageGenerationStatusExpired,
+ Prompt: "expired",
+ ModelName: "gpt-image-2",
+ RequestId: "req_expired",
+ FilePath: "",
+ }).Error)
+
+ items := GetAllUserDrawingLogs(1, 0, 10, TaskQueryParams{})
+ require.Len(t, items, 1)
+ require.Equal(t, "EXPIRED", items[0].Status)
+ require.Equal(t, "", items[0].ImageUrl)
+ require.Equal(t, "图片已过期", items[0].FailReason)
+}
diff --git a/model/task_cas_test.go b/model/task_cas_test.go
index 189d07c64207..708993f4147a 100644
--- a/model/task_cas_test.go
+++ b/model/task_cas_test.go
@@ -36,6 +36,7 @@ func TestMain(m *testing.M) {
if err := db.AutoMigrate(
&Task{},
&Midjourney{},
+ &ImageGeneration{},
&User{},
&Token{},
&Log{},
@@ -55,6 +56,8 @@ func truncateTables(t *testing.T) {
t.Helper()
t.Cleanup(func() {
DB.Exec("DELETE FROM tasks")
+ DB.Exec("DELETE FROM midjourneys")
+ DB.Exec("DELETE FROM image_generations")
DB.Exec("DELETE FROM users")
DB.Exec("DELETE FROM tokens")
DB.Exec("DELETE FROM logs")
diff --git a/relay/image_handler.go b/relay/image_handler.go
index a4fee7d9e0a6..046a10e4833b 100644
--- a/relay/image_handler.go
+++ b/relay/image_handler.go
@@ -106,7 +106,11 @@ func ImageHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *type
}
}
+ originalWriter := c.Writer
+ responseCapture := &imageResponseCaptureWriter{ResponseWriter: originalWriter}
+ c.Writer = responseCapture
usage, newAPIError := adaptor.DoResponse(c, httpResp, info)
+ c.Writer = originalWriter
if newAPIError != nil {
// reset status code 重置状态码
service.ResetStatusCode(newAPIError, statusCodeMappingStr)
@@ -150,6 +154,26 @@ func ImageHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *type
logContent = append(logContent, fmt.Sprintf("生成数量 %d", imageN))
}
- service.PostTextConsumeQuota(c, info, usage.(*dto.Usage), logContent)
+ summary := service.PostTextConsumeQuota(c, info, usage.(*dto.Usage), logContent)
+ service.SaveImageGenerationResponse(c, info, request, responseCapture.Bytes(), summary.Quota)
return nil
}
+
+type imageResponseCaptureWriter struct {
+ gin.ResponseWriter
+ body bytes.Buffer
+}
+
+func (w *imageResponseCaptureWriter) Write(data []byte) (int, error) {
+ w.body.Write(data)
+ return w.ResponseWriter.Write(data)
+}
+
+func (w *imageResponseCaptureWriter) WriteString(data string) (int, error) {
+ w.body.WriteString(data)
+ return w.ResponseWriter.WriteString(data)
+}
+
+func (w *imageResponseCaptureWriter) Bytes() []byte {
+ return w.body.Bytes()
+}
diff --git a/router/api-router.go b/router/api-router.go
index 10d19e6e311c..bf474ac85eff 100644
--- a/router/api-router.go
+++ b/router/api-router.go
@@ -324,6 +324,9 @@ func SetApiRouter(router *gin.Engine) {
mjRoute.GET("/self", middleware.UserAuth(), controller.GetUserMidjourney)
mjRoute.GET("/", middleware.AdminAuth(), controller.GetAllMidjourney)
+ imageGenerationRoute := apiRouter.Group("/image-generations")
+ imageGenerationRoute.GET("/:id/content", middleware.UserAuth(), controller.GetImageGenerationContent)
+
taskRoute := apiRouter.Group("/task")
{
taskRoute.GET("/self", middleware.UserAuth(), controller.GetUserTask)
diff --git a/service/image_generation_storage.go b/service/image_generation_storage.go
new file mode 100644
index 000000000000..4cee8fb8d44a
--- /dev/null
+++ b/service/image_generation_storage.go
@@ -0,0 +1,184 @@
+package service
+
+import (
+ "context"
+ "encoding/base64"
+ "fmt"
+ "net/http"
+ "os"
+ "path/filepath"
+ "strings"
+ "time"
+
+ "github.com/QuantumNous/new-api/common"
+ "github.com/QuantumNous/new-api/dto"
+ "github.com/QuantumNous/new-api/logger"
+ "github.com/QuantumNous/new-api/model"
+ relaycommon "github.com/QuantumNous/new-api/relay/common"
+
+ "github.com/gin-gonic/gin"
+)
+
+const imageGenerationRetention = 7 * 24 * time.Hour
+
+func imageGenerationStorageDir() string {
+ if dir := strings.TrimSpace(os.Getenv("IMAGE_GENERATION_STORAGE_DIR")); dir != "" {
+ return dir
+ }
+ if info, err := os.Stat("/data"); err == nil && info.IsDir() {
+ return "/data/image-generations"
+ }
+ return "data/image-generations"
+}
+
+func imageGenerationFilePath(relativePath string) string {
+ cleanPath := filepath.Clean(relativePath)
+ if filepath.IsAbs(cleanPath) || cleanPath == ".." || strings.HasPrefix(cleanPath, ".."+string(os.PathSeparator)) {
+ return filepath.Join(imageGenerationStorageDir(), "_invalid")
+ }
+ return filepath.Join(imageGenerationStorageDir(), cleanPath)
+}
+
+func SaveImageGenerationResponse(c *gin.Context, info *relaycommon.RelayInfo, request *dto.ImageRequest, responseBody []byte, quota int) {
+ if len(responseBody) == 0 || request == nil || info == nil {
+ return
+ }
+
+ var imageResponse dto.ImageResponse
+ if err := common.Unmarshal(responseBody, &imageResponse); err != nil {
+ logger.LogWarn(c, "failed to parse image generation response for storage: "+err.Error())
+ return
+ }
+ if len(imageResponse.Data) == 0 {
+ return
+ }
+
+ now := time.Now()
+ requestID := c.GetString(common.RequestIdKey)
+ if requestID == "" {
+ requestID = common.GetUUID()
+ }
+ perImageQuota := quota
+ if len(imageResponse.Data) > 0 {
+ perImageQuota = quota / len(imageResponse.Data)
+ }
+
+ for index, item := range imageResponse.Data {
+ if strings.TrimSpace(item.B64Json) == "" {
+ continue
+ }
+ mimeType, ext, raw, err := decodeImageGenerationBase64(item.B64Json)
+ if err != nil {
+ logger.LogWarn(c, fmt.Sprintf("failed to decode image generation response image %d: %s", index, err.Error()))
+ continue
+ }
+
+ relativeDir := filepath.Join(now.Format("20060102"), fmt.Sprintf("user-%d", info.UserId))
+ filename := fmt.Sprintf("%s-%d.%s", requestID, index, ext)
+ relativePath := filepath.Join(relativeDir, filename)
+ absolutePath := imageGenerationFilePath(relativePath)
+ if err := os.MkdirAll(filepath.Dir(absolutePath), 0750); err != nil {
+ logger.LogError(c, "failed to create image generation storage dir: "+err.Error())
+ continue
+ }
+ if err := os.WriteFile(absolutePath, raw, 0600); err != nil {
+ logger.LogError(c, "failed to write image generation file: "+err.Error())
+ continue
+ }
+
+ recordQuota := perImageQuota
+ if index == len(imageResponse.Data)-1 {
+ recordQuota = quota - perImageQuota*(len(imageResponse.Data)-1)
+ }
+ record := &model.ImageGeneration{
+ UserId: info.UserId,
+ TokenId: info.TokenId,
+ ChannelId: info.ChannelId,
+ RequestId: requestID,
+ ImageIndex: index,
+ ModelName: info.OriginModelName,
+ Prompt: request.Prompt,
+ Size: request.Size,
+ Quality: request.Quality,
+ Quota: recordQuota,
+ FilePath: relativePath,
+ MimeType: mimeType,
+ Status: model.ImageGenerationStatusSuccess,
+ Group: info.UsingGroup,
+ CreatedAt: now.Unix(),
+ ExpireAt: now.Add(imageGenerationRetention).Unix(),
+ }
+ if err := model.InsertImageGeneration(record); err != nil {
+ logger.LogError(c, "failed to insert image generation record: "+err.Error())
+ _ = os.Remove(absolutePath)
+ }
+ }
+}
+
+func decodeImageGenerationBase64(data string) (mimeType string, ext string, raw []byte, err error) {
+ if commaIndex := strings.Index(data, ","); commaIndex >= 0 {
+ data = data[commaIndex+1:]
+ }
+ raw, err = base64.StdEncoding.DecodeString(strings.TrimSpace(data))
+ if err != nil {
+ return "", "", nil, err
+ }
+ mimeType = http.DetectContentType(raw)
+ switch mimeType {
+ case "image/png":
+ ext = "png"
+ case "image/jpeg":
+ ext = "jpg"
+ case "image/webp":
+ ext = "webp"
+ case "image/gif":
+ ext = "gif"
+ default:
+ if strings.HasPrefix(mimeType, "image/") {
+ ext = strings.TrimPrefix(mimeType, "image/")
+ } else {
+ mimeType = "image/png"
+ ext = "png"
+ }
+ }
+ return mimeType, ext, raw, nil
+}
+
+func StartImageGenerationCleanupTask() {
+ go func() {
+ ticker := time.NewTicker(6 * time.Hour)
+ defer ticker.Stop()
+ for {
+ CleanupExpiredImageGenerations()
+ <-ticker.C
+ }
+ }()
+}
+
+func CleanupExpiredImageGenerations() {
+ for {
+ records, err := model.GetExpiredImageGenerations(time.Now().Unix(), 100)
+ if err != nil {
+ logger.LogError(context.Background(), "failed to query expired image generations: "+err.Error())
+ return
+ }
+ if len(records) == 0 {
+ return
+ }
+ for _, record := range records {
+ if record.FilePath != "" {
+ _ = os.Remove(imageGenerationFilePath(record.FilePath))
+ }
+ if err := model.MarkImageGenerationExpired(record.Id); err != nil {
+ logger.LogError(context.Background(), "failed to mark image generation expired: "+err.Error())
+ }
+ }
+ }
+}
+
+func GetImageGenerationAbsolutePath(record *model.ImageGeneration) string {
+ if record == nil || record.FilePath == "" {
+ return ""
+ }
+ return imageGenerationFilePath(record.FilePath)
+}
diff --git a/service/image_generation_storage_test.go b/service/image_generation_storage_test.go
new file mode 100644
index 000000000000..69abe3c40c51
--- /dev/null
+++ b/service/image_generation_storage_test.go
@@ -0,0 +1,104 @@
+package service
+
+import (
+ "net/http/httptest"
+ "os"
+ "path/filepath"
+ "testing"
+ "time"
+
+ "github.com/QuantumNous/new-api/common"
+ "github.com/QuantumNous/new-api/dto"
+ "github.com/QuantumNous/new-api/model"
+ relaycommon "github.com/QuantumNous/new-api/relay/common"
+
+ "github.com/gin-gonic/gin"
+ "github.com/stretchr/testify/require"
+)
+
+func TestSaveImageGenerationResponseStoresFileAndRecord(t *testing.T) {
+ truncateServiceImageGenerationTables(t)
+
+ storageDir := t.TempDir()
+ t.Setenv("IMAGE_GENERATION_STORAGE_DIR", storageDir)
+
+ gin.SetMode(gin.TestMode)
+ recorder := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(recorder)
+ c.Set(common.RequestIdKey, "req_image_store")
+
+ responseBody, err := common.Marshal(dto.ImageResponse{
+ Data: []dto.ImageData{
+ {B64Json: "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mP8/x8AAwMCAO+/p9sAAAAASUVORK5CYII="},
+ },
+ })
+ require.NoError(t, err)
+
+ relayInfo := &relaycommon.RelayInfo{
+ UserId: 7,
+ TokenId: 8,
+ OriginModelName: "gemini-3.1-flash-image",
+ UsingGroup: "Image",
+ ChannelMeta: &relaycommon.ChannelMeta{
+ ChannelId: 9,
+ },
+ }
+ request := &dto.ImageRequest{
+ Prompt: "a red cube",
+ Size: "1024x1024",
+ Quality: "standard",
+ }
+
+ SaveImageGenerationResponse(c, relayInfo, request, responseBody, 50000)
+
+ var records []model.ImageGeneration
+ require.NoError(t, model.DB.Find(&records).Error)
+ require.Len(t, records, 1)
+ require.Equal(t, "req_image_store", records[0].RequestId)
+ require.Equal(t, "gemini-3.1-flash-image", records[0].ModelName)
+ require.Equal(t, "a red cube", records[0].Prompt)
+ require.Equal(t, "1024x1024", records[0].Size)
+ require.Equal(t, 50000, records[0].Quota)
+ require.Equal(t, model.ImageGenerationStatusSuccess, records[0].Status)
+ require.NotEmpty(t, records[0].FilePath)
+
+ _, err = os.Stat(filepath.Join(storageDir, records[0].FilePath))
+ require.NoError(t, err)
+}
+
+func TestCleanupExpiredImageGenerationsDeletesFileAndMarksExpired(t *testing.T) {
+ truncateServiceImageGenerationTables(t)
+
+ storageDir := t.TempDir()
+ t.Setenv("IMAGE_GENERATION_STORAGE_DIR", storageDir)
+
+ relativePath := filepath.Join("20260710", "user-1", "expired.png")
+ absolutePath := filepath.Join(storageDir, relativePath)
+ require.NoError(t, os.MkdirAll(filepath.Dir(absolutePath), 0750))
+ require.NoError(t, os.WriteFile(absolutePath, []byte("png"), 0600))
+
+ record := &model.ImageGeneration{
+ UserId: 1,
+ RequestId: "req_expired",
+ FilePath: relativePath,
+ Status: model.ImageGenerationStatusSuccess,
+ CreatedAt: time.Now().Add(-8 * 24 * time.Hour).Unix(),
+ ExpireAt: time.Now().Add(-time.Hour).Unix(),
+ }
+ require.NoError(t, model.DB.Create(record).Error)
+
+ CleanupExpiredImageGenerations()
+
+ _, err := os.Stat(absolutePath)
+ require.True(t, os.IsNotExist(err))
+
+ var reloaded model.ImageGeneration
+ require.NoError(t, model.DB.First(&reloaded, record.Id).Error)
+ require.Equal(t, model.ImageGenerationStatusExpired, reloaded.Status)
+ require.Empty(t, reloaded.FilePath)
+}
+
+func truncateServiceImageGenerationTables(t *testing.T) {
+ t.Helper()
+ require.NoError(t, model.DB.Exec("DELETE FROM image_generations").Error)
+}
diff --git a/service/task_billing_test.go b/service/task_billing_test.go
index eb6b2a8444cd..981fcc9543ea 100644
--- a/service/task_billing_test.go
+++ b/service/task_billing_test.go
@@ -43,6 +43,7 @@ func TestMain(m *testing.M) {
&model.Token{},
&model.Log{},
&model.Channel{},
+ &model.ImageGeneration{},
&model.TopUp{},
&model.UserSubscription{},
); err != nil {
diff --git a/service/text_quota.go b/service/text_quota.go
index 8caee8f28799..c670e98367c1 100644
--- a/service/text_quota.go
+++ b/service/text_quota.go
@@ -291,7 +291,7 @@ func usageSemanticFromUsage(relayInfo *relaycommon.RelayInfo, usage *dto.Usage)
return "openai"
}
-func PostTextConsumeQuota(ctx *gin.Context, relayInfo *relaycommon.RelayInfo, usage *dto.Usage, extraContent []string) {
+func PostTextConsumeQuota(ctx *gin.Context, relayInfo *relaycommon.RelayInfo, usage *dto.Usage, extraContent []string) textQuotaSummary {
originUsage := usage
if usage == nil {
extraContent = append(extraContent, "上游无计费信息")
@@ -427,4 +427,5 @@ func PostTextConsumeQuota(ctx *gin.Context, relayInfo *relaycommon.RelayInfo, us
Group: relayInfo.UsingGroup,
Other: other,
})
+ return summary
}
diff --git a/web/src/components/table/mj-logs/MjLogsColumnDefs.jsx b/web/src/components/table/mj-logs/MjLogsColumnDefs.jsx
index c9ce662a3a9d..d9bf445b7e01 100644
--- a/web/src/components/table/mj-logs/MjLogsColumnDefs.jsx
+++ b/web/src/components/table/mj-logs/MjLogsColumnDefs.jsx
@@ -274,6 +274,12 @@ function renderStatus(type, t) {
{t('失败')}
);
+ case 'EXPIRED':
+ return (
+
}>
+ {t('已过期')}
+
+ );
case 'MODAL':
return (
Date: Fri, 10 Jul 2026 01:19:58 +0800
Subject: [PATCH 23/24] fix image generation log previews
---
controller/image_generation.go | 8 +--
controller/image_generation_test.go | 86 ++++++++++++++++++++++++
model/image_generation.go | 100 ++++++++++++++++++++++++++--
model/midjourney_test.go | 17 +++--
router/api-router.go | 2 +-
service/image_generation_storage.go | 8 +++
6 files changed, 206 insertions(+), 15 deletions(-)
create mode 100644 controller/image_generation_test.go
diff --git a/controller/image_generation.go b/controller/image_generation.go
index cc5f7118082b..05cb13010c62 100644
--- a/controller/image_generation.go
+++ b/controller/image_generation.go
@@ -5,7 +5,6 @@ import (
"os"
"strconv"
- "github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/model"
"github.com/QuantumNous/new-api/service"
@@ -25,10 +24,9 @@ func GetImageGenerationContent(c *gin.Context) {
return
}
- role := c.GetInt("role")
- userID := c.GetInt("id")
- if role < common.RoleAdminUser && record.UserId != userID {
- c.Status(http.StatusForbidden)
+ expires, err := strconv.ParseInt(c.Query("expires"), 10, 64)
+ if err != nil || !model.ValidateImageGenerationContentSignature(record, expires, c.Query("signature")) {
+ c.Status(http.StatusUnauthorized)
return
}
if record.Status != model.ImageGenerationStatusSuccess || record.FilePath == "" {
diff --git a/controller/image_generation_test.go b/controller/image_generation_test.go
new file mode 100644
index 000000000000..b16b1119d1e7
--- /dev/null
+++ b/controller/image_generation_test.go
@@ -0,0 +1,86 @@
+package controller
+
+import (
+ "fmt"
+ "net/http"
+ "net/http/httptest"
+ "os"
+ "path/filepath"
+ "testing"
+ "time"
+
+ "github.com/QuantumNous/new-api/common"
+ "github.com/QuantumNous/new-api/model"
+
+ "github.com/gin-gonic/gin"
+ "github.com/glebarez/sqlite"
+ "github.com/stretchr/testify/require"
+ "gorm.io/gorm"
+)
+
+func setupImageGenerationControllerTestDB(t *testing.T) *gorm.DB {
+ t.Helper()
+
+ gin.SetMode(gin.TestMode)
+ common.UsingSQLite = true
+ common.UsingMySQL = false
+ common.UsingPostgreSQL = false
+ common.RedisEnabled = false
+
+ db, err := gorm.Open(sqlite.Open(fmt.Sprintf("file:%s?mode=memory&cache=shared", t.Name())), &gorm.Config{})
+ require.NoError(t, err)
+ model.DB = db
+ model.LOG_DB = db
+ require.NoError(t, db.AutoMigrate(&model.ImageGeneration{}))
+
+ t.Cleanup(func() {
+ sqlDB, err := db.DB()
+ if err == nil {
+ _ = sqlDB.Close()
+ }
+ })
+ return db
+}
+
+func TestGetImageGenerationContentRequiresValidSignature(t *testing.T) {
+ db := setupImageGenerationControllerTestDB(t)
+
+ storageDir := t.TempDir()
+ t.Setenv("IMAGE_GENERATION_STORAGE_DIR", storageDir)
+
+ relativePath := filepath.Join("20260710", "user-1", "image.png")
+ absolutePath := filepath.Join(storageDir, relativePath)
+ require.NoError(t, os.MkdirAll(filepath.Dir(absolutePath), 0750))
+ require.NoError(t, os.WriteFile(absolutePath, []byte("png-data"), 0600))
+
+ record := &model.ImageGeneration{
+ UserId: 1,
+ RequestId: "req_image",
+ FilePath: relativePath,
+ MimeType: "image/png",
+ Status: model.ImageGenerationStatusSuccess,
+ CreatedAt: time.Now().Unix(),
+ ExpireAt: time.Now().Add(time.Hour).Unix(),
+ }
+ require.NoError(t, db.Create(record).Error)
+
+ router := gin.New()
+ router.GET("/api/image-generations/:id/content", GetImageGenerationContent)
+
+ missingSignature := httptest.NewRecorder()
+ req := httptest.NewRequest(http.MethodGet, fmt.Sprintf("/api/image-generations/%d/content", record.Id), nil)
+ router.ServeHTTP(missingSignature, req)
+ require.Equal(t, http.StatusUnauthorized, missingSignature.Code)
+
+ expires := record.ExpireAt
+ signature := model.GenerateImageGenerationContentSignature(record, expires)
+ valid := httptest.NewRecorder()
+ req = httptest.NewRequest(
+ http.MethodGet,
+ fmt.Sprintf("/api/image-generations/%d/content?expires=%d&signature=%s", record.Id, expires, signature),
+ nil,
+ )
+ router.ServeHTTP(valid, req)
+ require.Equal(t, http.StatusOK, valid.Code)
+ require.Equal(t, "png-data", valid.Body.String())
+}
diff --git a/model/image_generation.go b/model/image_generation.go
index 892c57e24aa5..1aad3a4c42c8 100644
--- a/model/image_generation.go
+++ b/model/image_generation.go
@@ -1,7 +1,12 @@
package model
import (
+ "fmt"
"strconv"
+ "strings"
+ "time"
+
+ "github.com/QuantumNous/new-api/common"
"gorm.io/gorm"
)
@@ -28,6 +33,7 @@ type ImageGeneration struct {
Status string `json:"status" gorm:"type:varchar(20);index"`
Group string `json:"group" gorm:"index"`
CreatedAt int64 `json:"created_at" gorm:"bigint;index"`
+ UseTime int64 `json:"use_time" gorm:"bigint"`
ExpireAt int64 `json:"expire_at" gorm:"bigint;index"`
}
@@ -113,7 +119,7 @@ func imageGenerationToMidjourney(record *ImageGeneration) *Midjourney {
status = ImageGenerationStatusSuccess
}
if status == ImageGenerationStatusSuccess && record.FilePath != "" {
- imageURL = "/api/image-generations/" + strconv.Itoa(record.Id) + "/content"
+ imageURL = imageGenerationContentURL(record)
}
if status == ImageGenerationStatusExpired {
failReason = "图片已过期"
@@ -123,6 +129,11 @@ func imageGenerationToMidjourney(record *ImageGeneration) *Midjourney {
if record.ImageIndex > 0 {
mjID = mjID + "#" + strconv.Itoa(record.ImageIndex+1)
}
+ useTime := imageGenerationUseTimeSeconds(record)
+ submitTime := record.CreatedAt * 1000
+ if useTime > 0 {
+ submitTime = (record.CreatedAt - useTime) * 1000
+ }
return &Midjourney{
Id: -record.Id,
@@ -130,10 +141,10 @@ func imageGenerationToMidjourney(record *ImageGeneration) *Midjourney {
UserId: record.UserId,
Action: "IMAGE_GENERATION",
MjId: mjID,
- Prompt: record.Prompt,
+ Prompt: imageGenerationPrompt(record),
PromptEn: record.ModelName,
- SubmitTime: record.CreatedAt * 1000,
- StartTime: record.CreatedAt * 1000,
+ SubmitTime: submitTime,
+ StartTime: submitTime,
FinishTime: record.CreatedAt * 1000,
ImageUrl: imageURL,
Status: status,
@@ -143,3 +154,84 @@ func imageGenerationToMidjourney(record *ImageGeneration) *Midjourney {
Quota: record.Quota,
}
}
+
+func imageGenerationPrompt(record *ImageGeneration) string {
+ parts := make([]string, 0, 4)
+ if record.Size != "" {
+ parts = append(parts, "大小 "+record.Size)
+ }
+ if record.Quality != "" {
+ parts = append(parts, "品质 "+record.Quality)
+ }
+ parts = append(parts, "生成数量 1")
+ if record.Prompt != "" {
+ parts = append(parts, "提示词 "+record.Prompt)
+ }
+ return strings.Join(parts, ", ")
+}
+
+func imageGenerationUseTimeSeconds(record *ImageGeneration) int64 {
+ if record.UseTime > 0 {
+ return record.UseTime
+ }
+ if record.RequestId == "" {
+ return 0
+ }
+ var log Log
+ result := LOG_DB.Model(&Log{}).
+ Select("use_time").
+ Where("request_id = ? AND type = ? AND use_time > 0", record.RequestId, LogTypeConsume).
+ Order("id desc").
+ Limit(1).
+ Find(&log)
+ if result.Error != nil || result.RowsAffected == 0 {
+ return 0
+ }
+ if log.UseTime < 0 {
+ return 0
+ }
+ return int64(log.UseTime)
+}
+
+func imageGenerationContentURL(record *ImageGeneration) string {
+ if record == nil {
+ return ""
+ }
+ expires := record.ExpireAt
+ if expires <= 0 {
+ expires = record.CreatedAt + 7*24*60*60
+ }
+ return fmt.Sprintf(
+ "/api/image-generations/%d/content?expires=%d&signature=%s",
+ record.Id,
+ expires,
+ GenerateImageGenerationContentSignature(record, expires),
+ )
+}
+
+func GenerateImageGenerationContentSignature(record *ImageGeneration, expires int64) string {
+ if record == nil {
+ return ""
+ }
+ payload := fmt.Sprintf(
+ "image-generation-content:%d:%d:%d:%s:%s:%d",
+ record.Id,
+ record.UserId,
+ expires,
+ record.FilePath,
+ record.Status,
+ record.ExpireAt,
+ )
+ return common.GenerateHMAC(payload)
+}
+
+func ValidateImageGenerationContentSignature(record *ImageGeneration, expires int64, signature string) bool {
+ if record == nil || signature == "" || expires <= time.Now().Unix() {
+ return false
+ }
+ if record.ExpireAt > 0 && expires > record.ExpireAt {
+ return false
+ }
+ expected := GenerateImageGenerationContentSignature(record, expires)
+ return expected != "" && expected == signature
+}
diff --git a/model/midjourney_test.go b/model/midjourney_test.go
index 3771f3be6c3b..76345dfedb39 100644
--- a/model/midjourney_test.go
+++ b/model/midjourney_test.go
@@ -26,9 +26,13 @@ func TestGetAllUserDrawingLogsIncludesImageGenerationLogs(t *testing.T) {
require.NoError(t, DB.Create(&ImageGeneration{
Id: 10,
UserId: 1,
- CreatedAt: 2,
+ CreatedAt: 100,
+ UseTime: 3,
+ ExpireAt: 4_102_444_800,
Status: ImageGenerationStatusSuccess,
Prompt: "a red cube",
+ Size: "1024x1024",
+ Quality: "standard",
ModelName: "gemini-3.1-flash-image",
Quota: 50000,
ChannelId: 23,
@@ -54,10 +58,13 @@ func TestGetAllUserDrawingLogsIncludesImageGenerationLogs(t *testing.T) {
require.Equal(t, "IMAGE_GENERATION", items[0].Action)
require.Equal(t, "SUCCESS", items[0].Status)
require.Equal(t, "100%", items[0].Progress)
- require.Equal(t, int64(2000), items[0].SubmitTime)
- require.Equal(t, int64(2000), items[0].FinishTime)
- require.Equal(t, "/api/image-generations/10/content", items[0].ImageUrl)
- require.Equal(t, "a red cube", items[0].Prompt)
+ require.Equal(t, int64(97000), items[0].SubmitTime)
+ require.Equal(t, int64(100000), items[0].FinishTime)
+ require.Contains(t, items[0].ImageUrl, "/api/image-generations/10/content?expires=")
+ require.Contains(t, items[0].ImageUrl, "signature=")
+ require.Contains(t, items[0].Prompt, "大小 1024x1024")
+ require.Contains(t, items[0].Prompt, "品质 standard")
+ require.Contains(t, items[0].Prompt, "提示词 a red cube")
require.Equal(t, "gemini-3.1-flash-image", items[0].PromptEn)
require.Equal(t, 50000, items[0].Quota)
require.Equal(t, "mj_old", items[1].MjId)
diff --git a/router/api-router.go b/router/api-router.go
index bf474ac85eff..b461518a258d 100644
--- a/router/api-router.go
+++ b/router/api-router.go
@@ -325,7 +325,7 @@ func SetApiRouter(router *gin.Engine) {
mjRoute.GET("/", middleware.AdminAuth(), controller.GetAllMidjourney)
imageGenerationRoute := apiRouter.Group("/image-generations")
- imageGenerationRoute.GET("/:id/content", middleware.UserAuth(), controller.GetImageGenerationContent)
+ imageGenerationRoute.GET("/:id/content", controller.GetImageGenerationContent)
taskRoute := apiRouter.Group("/task")
{
diff --git a/service/image_generation_storage.go b/service/image_generation_storage.go
index 4cee8fb8d44a..d0151f168560 100644
--- a/service/image_generation_storage.go
+++ b/service/image_generation_storage.go
@@ -54,6 +54,13 @@ func SaveImageGenerationResponse(c *gin.Context, info *relaycommon.RelayInfo, re
}
now := time.Now()
+ useTimeSeconds := int64(0)
+ if !info.StartTime.IsZero() {
+ useTimeSeconds = int64(now.Sub(info.StartTime).Seconds())
+ if useTimeSeconds < 0 {
+ useTimeSeconds = 0
+ }
+ }
requestID := c.GetString(common.RequestIdKey)
if requestID == "" {
requestID = common.GetUUID()
@@ -106,6 +113,7 @@ func SaveImageGenerationResponse(c *gin.Context, info *relaycommon.RelayInfo, re
Status: model.ImageGenerationStatusSuccess,
Group: info.UsingGroup,
CreatedAt: now.Unix(),
+ UseTime: useTimeSeconds,
ExpireAt: now.Add(imageGenerationRetention).Unix(),
}
if err := model.InsertImageGeneration(record); err != nil {
From 6c1ad65380759d3128a4820978dfebd01dd09705 Mon Sep 17 00:00:00 2001
From: "649985538@qq.com" <649985538@qq.com>
Date: Fri, 10 Jul 2026 01:29:18 +0800
Subject: [PATCH 24/24] default image generation quality label
---
model/image_generation.go | 8 ++++++--
service/image_generation_storage.go | 6 +++++-
2 files changed, 11 insertions(+), 3 deletions(-)
diff --git a/model/image_generation.go b/model/image_generation.go
index 1aad3a4c42c8..62d2d747a269 100644
--- a/model/image_generation.go
+++ b/model/image_generation.go
@@ -160,8 +160,12 @@ func imageGenerationPrompt(record *ImageGeneration) string {
if record.Size != "" {
parts = append(parts, "大小 "+record.Size)
}
- if record.Quality != "" {
- parts = append(parts, "品质 "+record.Quality)
+ quality := record.Quality
+ if quality == "" {
+ quality = "standard"
+ }
+ if quality != "" {
+ parts = append(parts, "品质 "+quality)
}
parts = append(parts, "生成数量 1")
if record.Prompt != "" {
diff --git a/service/image_generation_storage.go b/service/image_generation_storage.go
index d0151f168560..b51ac76eb820 100644
--- a/service/image_generation_storage.go
+++ b/service/image_generation_storage.go
@@ -97,6 +97,10 @@ func SaveImageGenerationResponse(c *gin.Context, info *relaycommon.RelayInfo, re
if index == len(imageResponse.Data)-1 {
recordQuota = quota - perImageQuota*(len(imageResponse.Data)-1)
}
+ quality := request.Quality
+ if quality == "" {
+ quality = "standard"
+ }
record := &model.ImageGeneration{
UserId: info.UserId,
TokenId: info.TokenId,
@@ -106,7 +110,7 @@ func SaveImageGenerationResponse(c *gin.Context, info *relaycommon.RelayInfo, re
ModelName: info.OriginModelName,
Prompt: request.Prompt,
Size: request.Size,
- Quality: request.Quality,
+ Quality: quality,
Quota: recordQuota,
FilePath: relativePath,
MimeType: mimeType,