Skip to content
Toggle navigation
P
Projects
G
Groups
S
Snippets
Help
phsl
/
new-api
This project
Loading...
Sign in
Toggle navigation
Go to a project
Project
Repository
Issues
0
Merge Requests
0
Pipelines
Wiki
Snippets
Members
Activity
Graph
Charts
Create a new issue
Jobs
Commits
Issue Boards
Files
Commits
Branches
Tags
Contributors
Graph
Compare
Charts
Commit
9d685276
authored
Nov 15, 2023
by
CaIon
Browse files
Options
Browse Files
Download
Email Patches
Plain Diff
support tts
parent
28fe74d7
Show whitespace changes
Inline
Side-by-side
Showing
7 changed files
with
76 additions
and
8 deletions
+76
-8
common/model-ratio.go
+2
-0
common/utils.go
+9
-0
controller/relay-audio.go
+44
-6
controller/relay-utils.go
+9
-0
controller/relay.go
+6
-0
middleware/distributor.go
+5
-2
router/relay-router.go
+1
-0
No files found.
common/model-ratio.go
View file @
9d685276
...
@@ -38,6 +38,8 @@ var ModelRatio = map[string]float64{
...
@@ -38,6 +38,8 @@ var ModelRatio = map[string]float64{
"text-davinci-edit-001"
:
10
,
"text-davinci-edit-001"
:
10
,
"code-davinci-edit-001"
:
10
,
"code-davinci-edit-001"
:
10
,
"whisper-1"
:
15
,
// $0.006 / minute -> $0.006 / 150 words -> $0.006 / 200 tokens -> $0.03 / 1k tokens
"whisper-1"
:
15
,
// $0.006 / minute -> $0.006 / 150 words -> $0.006 / 200 tokens -> $0.03 / 1k tokens
"tts-1"
:
7.5
,
// 1k characters -> $0.015
"tts-1-hd"
:
15
,
// 1k characters -> $0.03
"davinci"
:
10
,
"davinci"
:
10
,
"curie"
:
10
,
"curie"
:
10
,
"babbage"
:
10
,
"babbage"
:
10
,
...
...
common/utils.go
View file @
9d685276
...
@@ -207,3 +207,12 @@ func String2Int(str string) int {
...
@@ -207,3 +207,12 @@ func String2Int(str string) int {
}
}
return
num
return
num
}
}
func
StringsContains
(
strs
[]
string
,
str
string
)
bool
{
for
_
,
s
:=
range
strs
{
if
s
==
str
{
return
true
}
}
return
false
}
controller/relay-audio.go
View file @
9d685276
...
@@ -11,10 +11,19 @@ import (
...
@@ -11,10 +11,19 @@ import (
"net/http"
"net/http"
"one-api/common"
"one-api/common"
"one-api/model"
"one-api/model"
"strings"
)
)
var
availableVoices
=
[]
string
{
"alloy"
,
"echo"
,
"fable"
,
"onyx"
,
"nova"
,
"shimmer"
,
}
func
relayAudioHelper
(
c
*
gin
.
Context
,
relayMode
int
)
*
OpenAIErrorWithStatusCode
{
func
relayAudioHelper
(
c
*
gin
.
Context
,
relayMode
int
)
*
OpenAIErrorWithStatusCode
{
audioModel
:=
"whisper-1"
tokenId
:=
c
.
GetInt
(
"token_id"
)
tokenId
:=
c
.
GetInt
(
"token_id"
)
channelType
:=
c
.
GetInt
(
"channel"
)
channelType
:=
c
.
GetInt
(
"channel"
)
...
@@ -22,8 +31,28 @@ func relayAudioHelper(c *gin.Context, relayMode int) *OpenAIErrorWithStatusCode
...
@@ -22,8 +31,28 @@ func relayAudioHelper(c *gin.Context, relayMode int) *OpenAIErrorWithStatusCode
userId
:=
c
.
GetInt
(
"id"
)
userId
:=
c
.
GetInt
(
"id"
)
group
:=
c
.
GetString
(
"group"
)
group
:=
c
.
GetString
(
"group"
)
var
audioRequest
AudioRequest
err
:=
common
.
UnmarshalBodyReusable
(
c
,
&
audioRequest
)
if
err
!=
nil
{
return
errorWrapper
(
err
,
"bind_request_body_failed"
,
http
.
StatusBadRequest
)
}
// request validation
if
audioRequest
.
Model
==
""
{
return
errorWrapper
(
errors
.
New
(
"model is required"
),
"required_field_missing"
,
http
.
StatusBadRequest
)
}
if
strings
.
HasPrefix
(
audioRequest
.
Model
,
"tts-1"
)
{
if
audioRequest
.
Voice
==
""
{
return
errorWrapper
(
errors
.
New
(
"voice is required"
),
"required_field_missing"
,
http
.
StatusBadRequest
)
}
if
!
common
.
StringsContains
(
availableVoices
,
audioRequest
.
Voice
)
{
return
errorWrapper
(
errors
.
New
(
"voice must be one of "
+
strings
.
Join
(
availableVoices
,
", "
)),
"invalid_field_value"
,
http
.
StatusBadRequest
)
}
}
preConsumedTokens
:=
common
.
PreConsumedQuota
preConsumedTokens
:=
common
.
PreConsumedQuota
modelRatio
:=
common
.
GetModelRatio
(
audioModel
)
modelRatio
:=
common
.
GetModelRatio
(
audio
Request
.
Model
)
groupRatio
:=
common
.
GetGroupRatio
(
group
)
groupRatio
:=
common
.
GetGroupRatio
(
group
)
ratio
:=
modelRatio
*
groupRatio
ratio
:=
modelRatio
*
groupRatio
preConsumedQuota
:=
int
(
float64
(
preConsumedTokens
)
*
ratio
)
preConsumedQuota
:=
int
(
float64
(
preConsumedTokens
)
*
ratio
)
...
@@ -58,8 +87,8 @@ func relayAudioHelper(c *gin.Context, relayMode int) *OpenAIErrorWithStatusCode
...
@@ -58,8 +87,8 @@ func relayAudioHelper(c *gin.Context, relayMode int) *OpenAIErrorWithStatusCode
if
err
!=
nil
{
if
err
!=
nil
{
return
errorWrapper
(
err
,
"unmarshal_model_mapping_failed"
,
http
.
StatusInternalServerError
)
return
errorWrapper
(
err
,
"unmarshal_model_mapping_failed"
,
http
.
StatusInternalServerError
)
}
}
if
modelMap
[
audioModel
]
!=
""
{
if
modelMap
[
audio
Request
.
Model
]
!=
""
{
audio
Model
=
modelMap
[
audio
Model
]
audio
Request
.
Model
=
modelMap
[
audioRequest
.
Model
]
}
}
}
}
...
@@ -97,7 +126,12 @@ func relayAudioHelper(c *gin.Context, relayMode int) *OpenAIErrorWithStatusCode
...
@@ -97,7 +126,12 @@ func relayAudioHelper(c *gin.Context, relayMode int) *OpenAIErrorWithStatusCode
defer
func
(
ctx
context
.
Context
)
{
defer
func
(
ctx
context
.
Context
)
{
go
func
()
{
go
func
()
{
quota
:=
countTokenText
(
audioResponse
.
Text
,
audioModel
)
var
quota
int
if
strings
.
HasPrefix
(
audioRequest
.
Model
,
"tts-1"
)
{
quota
=
countAudioToken
(
audioRequest
.
Input
,
audioRequest
.
Model
)
}
else
{
quota
=
countAudioToken
(
audioResponse
.
Text
,
audioRequest
.
Model
)
}
quotaDelta
:=
quota
-
preConsumedQuota
quotaDelta
:=
quota
-
preConsumedQuota
err
:=
model
.
PostConsumeTokenQuota
(
tokenId
,
userQuota
,
quotaDelta
,
preConsumedQuota
,
true
)
err
:=
model
.
PostConsumeTokenQuota
(
tokenId
,
userQuota
,
quotaDelta
,
preConsumedQuota
,
true
)
if
err
!=
nil
{
if
err
!=
nil
{
...
@@ -110,7 +144,7 @@ func relayAudioHelper(c *gin.Context, relayMode int) *OpenAIErrorWithStatusCode
...
@@ -110,7 +144,7 @@ func relayAudioHelper(c *gin.Context, relayMode int) *OpenAIErrorWithStatusCode
if
quota
!=
0
{
if
quota
!=
0
{
tokenName
:=
c
.
GetString
(
"token_name"
)
tokenName
:=
c
.
GetString
(
"token_name"
)
logContent
:=
fmt
.
Sprintf
(
"模型倍率 %.2f,分组倍率 %.2f"
,
modelRatio
,
groupRatio
)
logContent
:=
fmt
.
Sprintf
(
"模型倍率 %.2f,分组倍率 %.2f"
,
modelRatio
,
groupRatio
)
model
.
RecordConsumeLog
(
ctx
,
userId
,
channelId
,
0
,
0
,
audioModel
,
tokenName
,
quota
,
logContent
,
tokenId
)
model
.
RecordConsumeLog
(
ctx
,
userId
,
channelId
,
0
,
0
,
audio
Request
.
Model
,
tokenName
,
quota
,
logContent
,
tokenId
)
model
.
UpdateUserUsedQuotaAndRequestCount
(
userId
,
quota
)
model
.
UpdateUserUsedQuotaAndRequestCount
(
userId
,
quota
)
channelId
:=
c
.
GetInt
(
"channel_id"
)
channelId
:=
c
.
GetInt
(
"channel_id"
)
model
.
UpdateChannelUsedQuota
(
channelId
,
quota
)
model
.
UpdateChannelUsedQuota
(
channelId
,
quota
)
...
@@ -127,10 +161,14 @@ func relayAudioHelper(c *gin.Context, relayMode int) *OpenAIErrorWithStatusCode
...
@@ -127,10 +161,14 @@ func relayAudioHelper(c *gin.Context, relayMode int) *OpenAIErrorWithStatusCode
if
err
!=
nil
{
if
err
!=
nil
{
return
errorWrapper
(
err
,
"close_response_body_failed"
,
http
.
StatusInternalServerError
)
return
errorWrapper
(
err
,
"close_response_body_failed"
,
http
.
StatusInternalServerError
)
}
}
if
strings
.
HasPrefix
(
audioRequest
.
Model
,
"tts-1"
)
{
}
else
{
err
=
json
.
Unmarshal
(
responseBody
,
&
audioResponse
)
err
=
json
.
Unmarshal
(
responseBody
,
&
audioResponse
)
if
err
!=
nil
{
if
err
!=
nil
{
return
errorWrapper
(
err
,
"unmarshal_response_body_failed"
,
http
.
StatusInternalServerError
)
return
errorWrapper
(
err
,
"unmarshal_response_body_failed"
,
http
.
StatusInternalServerError
)
}
}
}
resp
.
Body
=
io
.
NopCloser
(
bytes
.
NewBuffer
(
responseBody
))
resp
.
Body
=
io
.
NopCloser
(
bytes
.
NewBuffer
(
responseBody
))
...
...
controller/relay-utils.go
View file @
9d685276
...
@@ -10,6 +10,7 @@ import (
...
@@ -10,6 +10,7 @@ import (
"one-api/common"
"one-api/common"
"strconv"
"strconv"
"strings"
"strings"
"unicode/utf8"
)
)
var
stopFinishReason
=
"stop"
var
stopFinishReason
=
"stop"
...
@@ -106,6 +107,14 @@ func countTokenInput(input any, model string) int {
...
@@ -106,6 +107,14 @@ func countTokenInput(input any, model string) int {
return
0
return
0
}
}
func
countAudioToken
(
text
string
,
model
string
)
int
{
if
strings
.
HasPrefix
(
model
,
"tts"
)
{
return
utf8
.
RuneCountInString
(
text
)
}
else
{
return
countTokenText
(
text
,
model
)
}
}
func
countTokenText
(
text
string
,
model
string
)
int
{
func
countTokenText
(
text
string
,
model
string
)
int
{
tokenEncoder
:=
getTokenEncoder
(
model
)
tokenEncoder
:=
getTokenEncoder
(
model
)
return
getTokenNum
(
tokenEncoder
,
text
)
return
getTokenNum
(
tokenEncoder
,
text
)
...
...
controller/relay.go
View file @
9d685276
...
@@ -70,6 +70,12 @@ func (r GeneralOpenAIRequest) ParseInput() []string {
...
@@ -70,6 +70,12 @@ func (r GeneralOpenAIRequest) ParseInput() []string {
return
input
return
input
}
}
type
AudioRequest
struct
{
Model
string
`json:"model"`
Voice
string
`json:"voice"`
Input
string
`json:"input"`
}
type
ChatRequest
struct
{
type
ChatRequest
struct
{
Model
string
`json:"model"`
Model
string
`json:"model"`
Messages
[]
Message
`json:"messages"`
Messages
[]
Message
`json:"messages"`
...
...
middleware/distributor.go
View file @
9d685276
...
@@ -46,9 +46,8 @@ func Distribute() func(c *gin.Context) {
...
@@ -46,9 +46,8 @@ func Distribute() func(c *gin.Context) {
if
modelRequest
.
Model
==
""
{
if
modelRequest
.
Model
==
""
{
modelRequest
.
Model
=
"midjourney"
modelRequest
.
Model
=
"midjourney"
}
}
}
else
if
!
strings
.
HasPrefix
(
c
.
Request
.
URL
.
Path
,
"/v1/audio"
)
{
err
=
common
.
UnmarshalBodyReusable
(
c
,
&
modelRequest
)
}
}
err
=
common
.
UnmarshalBodyReusable
(
c
,
&
modelRequest
)
if
err
!=
nil
{
if
err
!=
nil
{
abortWithMessage
(
c
,
http
.
StatusBadRequest
,
"无效的请求"
)
abortWithMessage
(
c
,
http
.
StatusBadRequest
,
"无效的请求"
)
return
return
...
@@ -70,9 +69,13 @@ func Distribute() func(c *gin.Context) {
...
@@ -70,9 +69,13 @@ func Distribute() func(c *gin.Context) {
}
}
if
strings
.
HasPrefix
(
c
.
Request
.
URL
.
Path
,
"/v1/audio"
)
{
if
strings
.
HasPrefix
(
c
.
Request
.
URL
.
Path
,
"/v1/audio"
)
{
if
modelRequest
.
Model
==
""
{
if
modelRequest
.
Model
==
""
{
if
strings
.
HasPrefix
(
c
.
Request
.
URL
.
Path
,
"/v1/audio/speech"
)
{
modelRequest
.
Model
=
"tts-1"
}
else
{
modelRequest
.
Model
=
"whisper-1"
modelRequest
.
Model
=
"whisper-1"
}
}
}
}
}
channel
,
err
=
model
.
CacheGetRandomSatisfiedChannel
(
userGroup
,
modelRequest
.
Model
)
channel
,
err
=
model
.
CacheGetRandomSatisfiedChannel
(
userGroup
,
modelRequest
.
Model
)
if
err
!=
nil
{
if
err
!=
nil
{
message
:=
fmt
.
Sprintf
(
"当前分组 %s 下对于模型 %s 无可用渠道"
,
userGroup
,
modelRequest
.
Model
)
message
:=
fmt
.
Sprintf
(
"当前分组 %s 下对于模型 %s 无可用渠道"
,
userGroup
,
modelRequest
.
Model
)
...
...
router/relay-router.go
View file @
9d685276
...
@@ -29,6 +29,7 @@ func SetRelayRouter(router *gin.Engine) {
...
@@ -29,6 +29,7 @@ func SetRelayRouter(router *gin.Engine) {
relayV1Router
.
POST
(
"/engines/:model/embeddings"
,
controller
.
Relay
)
relayV1Router
.
POST
(
"/engines/:model/embeddings"
,
controller
.
Relay
)
relayV1Router
.
POST
(
"/audio/transcriptions"
,
controller
.
Relay
)
relayV1Router
.
POST
(
"/audio/transcriptions"
,
controller
.
Relay
)
relayV1Router
.
POST
(
"/audio/translations"
,
controller
.
Relay
)
relayV1Router
.
POST
(
"/audio/translations"
,
controller
.
Relay
)
relayV1Router
.
POST
(
"/audio/speech"
,
controller
.
Relay
)
relayV1Router
.
GET
(
"/files"
,
controller
.
RelayNotImplemented
)
relayV1Router
.
GET
(
"/files"
,
controller
.
RelayNotImplemented
)
relayV1Router
.
POST
(
"/files"
,
controller
.
RelayNotImplemented
)
relayV1Router
.
POST
(
"/files"
,
controller
.
RelayNotImplemented
)
relayV1Router
.
DELETE
(
"/files/:id"
,
controller
.
RelayNotImplemented
)
relayV1Router
.
DELETE
(
"/files/:id"
,
controller
.
RelayNotImplemented
)
...
...
Write
Preview
Markdown
is supported
0%
Try again
or
attach a new file
Attach a file
Cancel
You are about to add
0
people
to the discussion. Proceed with caution.
Finish editing this message first!
Cancel
Please
register
or
sign in
to comment