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
a4cf25f4
authored
Aug 31, 2023
by
CaIon
Browse files
Options
Browse Files
Download
Email Patches
Plain Diff
修复多路复用bug
parent
46d817ab
Hide whitespace changes
Inline
Side-by-side
Showing
5 changed files
with
42 additions
and
8 deletions
+42
-8
controller/midjourney.go
+4
-0
controller/relay-mj.go
+5
-1
controller/relay-text.go
+18
-1
controller/relay.go
+2
-2
middleware/auth.go
+13
-4
No files found.
controller/midjourney.go
View file @
a4cf25f4
...
@@ -37,22 +37,26 @@ func UpdateMidjourneyTask() {
...
@@ -37,22 +37,26 @@ func UpdateMidjourneyTask() {
jsonStr
,
err
:=
json
.
Marshal
(
requestBody
)
jsonStr
,
err
:=
json
.
Marshal
(
requestBody
)
if
err
!=
nil
{
if
err
!=
nil
{
log
.
Printf
(
"UpdateMidjourneyTask: %v"
,
err
)
log
.
Printf
(
"UpdateMidjourneyTask: %v"
,
err
)
continue
}
}
req
,
err
:=
http
.
NewRequest
(
"POST"
,
requestUrl
,
bytes
.
NewBuffer
(
jsonStr
))
req
,
err
:=
http
.
NewRequest
(
"POST"
,
requestUrl
,
bytes
.
NewBuffer
(
jsonStr
))
if
err
!=
nil
{
if
err
!=
nil
{
log
.
Printf
(
"UpdateMidjourneyTask: %v"
,
err
)
log
.
Printf
(
"UpdateMidjourneyTask: %v"
,
err
)
continue
}
}
req
.
Header
.
Set
(
"Content-Type"
,
"application/json"
)
req
.
Header
.
Set
(
"Content-Type"
,
"application/json"
)
req
.
Header
.
Set
(
"mj-api-secret"
,
"uhiftyuwadbkjshbiklahcuitguasguzhxliawodawdu"
)
req
.
Header
.
Set
(
"mj-api-secret"
,
"uhiftyuwadbkjshbiklahcuitguasguzhxliawodawdu"
)
resp
,
err
:=
httpClient
.
Do
(
req
)
resp
,
err
:=
httpClient
.
Do
(
req
)
if
err
!=
nil
{
if
err
!=
nil
{
log
.
Printf
(
"UpdateMidjourneyTask: %v"
,
err
)
log
.
Printf
(
"UpdateMidjourneyTask: %v"
,
err
)
continue
}
}
defer
resp
.
Body
.
Close
()
defer
resp
.
Body
.
Close
()
var
response
[]
Midjourney
var
response
[]
Midjourney
err
=
json
.
NewDecoder
(
resp
.
Body
)
.
Decode
(
&
response
)
err
=
json
.
NewDecoder
(
resp
.
Body
)
.
Decode
(
&
response
)
if
err
!=
nil
{
if
err
!=
nil
{
log
.
Printf
(
"UpdateMidjourneyTask: %v"
,
err
)
log
.
Printf
(
"UpdateMidjourneyTask: %v"
,
err
)
continue
}
}
for
_
,
responseItem
:=
range
response
{
for
_
,
responseItem
:=
range
response
{
var
midjourneyTask
*
model
.
Midjourney
var
midjourneyTask
*
model
.
Midjourney
...
...
controller/relay-mj.go
View file @
a4cf25f4
...
@@ -248,6 +248,10 @@ func relayMidjourneySubmit(c *gin.Context, relayMode int) *MidjourneyResponse {
...
@@ -248,6 +248,10 @@ func relayMidjourneySubmit(c *gin.Context, relayMode int) *MidjourneyResponse {
req
.
Header
.
Set
(
"Content-Type"
,
c
.
Request
.
Header
.
Get
(
"Content-Type"
))
req
.
Header
.
Set
(
"Content-Type"
,
c
.
Request
.
Header
.
Get
(
"Content-Type"
))
req
.
Header
.
Set
(
"Accept"
,
c
.
Request
.
Header
.
Get
(
"Accept"
))
req
.
Header
.
Set
(
"Accept"
,
c
.
Request
.
Header
.
Get
(
"Accept"
))
//mjToken := ""
//if c.Request.Header.Get("Authorization") != "" {
// mjToken = strings.Split(c.Request.Header.Get("Authorization"), " ")[1]
//}
req
.
Header
.
Set
(
"mj-api-secret"
,
strings
.
Split
(
c
.
Request
.
Header
.
Get
(
"Authorization"
),
" "
)[
1
])
req
.
Header
.
Set
(
"mj-api-secret"
,
strings
.
Split
(
c
.
Request
.
Header
.
Get
(
"Authorization"
),
" "
)[
1
])
// print request header
// print request header
log
.
Printf
(
"request header: %s"
,
req
.
Header
)
log
.
Printf
(
"request header: %s"
,
req
.
Header
)
...
@@ -353,7 +357,7 @@ func relayMidjourneySubmit(c *gin.Context, relayMode int) *MidjourneyResponse {
...
@@ -353,7 +357,7 @@ func relayMidjourneySubmit(c *gin.Context, relayMode int) *MidjourneyResponse {
Progress
:
"0%"
,
Progress
:
"0%"
,
FailReason
:
""
,
FailReason
:
""
,
}
}
if
midjResponse
.
Code
==
4
{
if
midjResponse
.
Code
==
4
||
midjResponse
.
Code
==
24
{
midjourneyTask
.
FailReason
=
midjResponse
.
Description
midjourneyTask
.
FailReason
=
midjResponse
.
Description
}
}
err
=
midjourneyTask
.
Insert
()
err
=
midjourneyTask
.
Insert
()
...
...
controller/relay-text.go
View file @
a4cf25f4
...
@@ -7,6 +7,7 @@ import (
...
@@ -7,6 +7,7 @@ import (
"fmt"
"fmt"
"github.com/gin-gonic/gin"
"github.com/gin-gonic/gin"
"io"
"io"
"log"
"net/http"
"net/http"
"one-api/common"
"one-api/common"
"one-api/model"
"one-api/model"
...
@@ -278,6 +279,10 @@ func relayTextHelper(c *gin.Context, relayMode int) *OpenAIErrorWithStatusCode {
...
@@ -278,6 +279,10 @@ func relayTextHelper(c *gin.Context, relayMode int) *OpenAIErrorWithStatusCode {
if
apiType
!=
APITypeXunfei
{
// cause xunfei use websocket
if
apiType
!=
APITypeXunfei
{
// cause xunfei use websocket
req
,
err
=
http
.
NewRequest
(
c
.
Request
.
Method
,
fullRequestURL
,
requestBody
)
req
,
err
=
http
.
NewRequest
(
c
.
Request
.
Method
,
fullRequestURL
,
requestBody
)
// 设置GetBody函数,该函数返回一个新的io.ReadCloser,该io.ReadCloser返回与原始请求体相同的数据
req
.
GetBody
=
func
()
(
io
.
ReadCloser
,
error
)
{
return
io
.
NopCloser
(
requestBody
),
nil
}
if
err
!=
nil
{
if
err
!=
nil
{
return
errorWrapper
(
err
,
"new_request_failed"
,
http
.
StatusInternalServerError
)
return
errorWrapper
(
err
,
"new_request_failed"
,
http
.
StatusInternalServerError
)
}
}
...
@@ -308,7 +313,9 @@ func relayTextHelper(c *gin.Context, relayMode int) *OpenAIErrorWithStatusCode {
...
@@ -308,7 +313,9 @@ func relayTextHelper(c *gin.Context, relayMode int) *OpenAIErrorWithStatusCode {
}
}
req
.
Header
.
Set
(
"Content-Type"
,
c
.
Request
.
Header
.
Get
(
"Content-Type"
))
req
.
Header
.
Set
(
"Content-Type"
,
c
.
Request
.
Header
.
Get
(
"Content-Type"
))
req
.
Header
.
Set
(
"Accept"
,
c
.
Request
.
Header
.
Get
(
"Accept"
))
req
.
Header
.
Set
(
"Accept"
,
c
.
Request
.
Header
.
Get
(
"Accept"
))
//req.Header.Set("Connection", c.Request.Header.Get("Connection"))
//req.Header.Set("Connection", c.Request.Header.Get("Connection"))
req
.
Close
=
true
resp
,
err
=
httpClient
.
Do
(
req
)
resp
,
err
=
httpClient
.
Do
(
req
)
if
err
!=
nil
{
if
err
!=
nil
{
return
errorWrapper
(
err
,
"do_request_failed"
,
http
.
StatusInternalServerError
)
return
errorWrapper
(
err
,
"do_request_failed"
,
http
.
StatusInternalServerError
)
...
@@ -324,8 +331,18 @@ func relayTextHelper(c *gin.Context, relayMode int) *OpenAIErrorWithStatusCode {
...
@@ -324,8 +331,18 @@ func relayTextHelper(c *gin.Context, relayMode int) *OpenAIErrorWithStatusCode {
isStream
=
isStream
||
strings
.
HasPrefix
(
resp
.
Header
.
Get
(
"Content-Type"
),
"text/event-stream"
)
isStream
=
isStream
||
strings
.
HasPrefix
(
resp
.
Header
.
Get
(
"Content-Type"
),
"text/event-stream"
)
if
resp
.
StatusCode
!=
http
.
StatusOK
{
if
resp
.
StatusCode
!=
http
.
StatusOK
{
//print resp body
body
,
err
:=
io
.
ReadAll
(
resp
.
Body
)
if
err
!=
nil
{
log
.
Println
(
"read resp err body failed"
,
err
)
}
log
.
Println
(
"resp body:"
,
string
(
body
))
errStr
:=
fmt
.
Sprintf
(
"bad status code: %d"
,
resp
.
StatusCode
)
if
resp
.
StatusCode
==
503
{
errStr
=
string
(
body
)
}
return
errorWrapper
(
return
errorWrapper
(
fmt
.
Errorf
(
"bad status code: %d"
,
resp
.
StatusCode
),
"bad_status_code"
,
resp
.
StatusCode
)
fmt
.
Errorf
(
errStr
),
"bad_status_code"
,
resp
.
StatusCode
)
}
}
}
}
...
...
controller/relay.go
View file @
a4cf25f4
...
@@ -205,7 +205,7 @@ func Relay(c *gin.Context) {
...
@@ -205,7 +205,7 @@ func Relay(c *gin.Context) {
})
})
}
}
channelId
:=
c
.
GetInt
(
"channel_id"
)
channelId
:=
c
.
GetInt
(
"channel_id"
)
common
.
SysError
(
fmt
.
Sprintf
(
"relay error (channel #%d): %
s"
,
channelId
,
err
.
Message
))
common
.
SysError
(
fmt
.
Sprintf
(
"relay error (channel #%d): %
v "
,
channelId
,
err
))
// https://platform.openai.com/docs/guides/error-codes/api-errors
// https://platform.openai.com/docs/guides/error-codes/api-errors
if
shouldDisableChannel
(
&
err
.
OpenAIError
)
{
if
shouldDisableChannel
(
&
err
.
OpenAIError
)
{
channelId
:=
c
.
GetInt
(
"channel_id"
)
channelId
:=
c
.
GetInt
(
"channel_id"
)
...
@@ -259,7 +259,7 @@ func RelayMidjourney(c *gin.Context) {
...
@@ -259,7 +259,7 @@ func RelayMidjourney(c *gin.Context) {
// channelId := c.GetInt("channel_id")
// channelId := c.GetInt("channel_id")
// channelName := c.GetString("channel_name")
// channelName := c.GetString("channel_name")
// disableChannel(channelId, channelName, err.Result)
// disableChannel(channelId, channelName, err.Result)
//}
//}
;''''''''''''''''''''''''''''''''
}
}
}
}
...
...
middleware/auth.go
View file @
a4cf25f4
...
@@ -85,10 +85,19 @@ func RootAuth() func(c *gin.Context) {
...
@@ -85,10 +85,19 @@ func RootAuth() func(c *gin.Context) {
func
TokenAuth
()
func
(
c
*
gin
.
Context
)
{
func
TokenAuth
()
func
(
c
*
gin
.
Context
)
{
return
func
(
c
*
gin
.
Context
)
{
return
func
(
c
*
gin
.
Context
)
{
key
:=
c
.
Request
.
Header
.
Get
(
"Authorization"
)
key
:=
c
.
Request
.
Header
.
Get
(
"Authorization"
)
key
=
strings
.
TrimPrefix
(
key
,
"Bearer "
)
parts
:=
make
([]
string
,
0
)
key
=
strings
.
TrimPrefix
(
key
,
"sk-"
)
if
key
==
""
{
parts
:=
strings
.
Split
(
key
,
"-"
)
key
=
c
.
Request
.
Header
.
Get
(
"mj-api-secret"
)
key
=
parts
[
0
]
key
=
strings
.
TrimPrefix
(
key
,
"Bearer "
)
key
=
strings
.
TrimPrefix
(
key
,
"sk-"
)
parts
:=
strings
.
Split
(
key
,
"-"
)
key
=
parts
[
0
]
}
else
{
key
=
strings
.
TrimPrefix
(
key
,
"Bearer "
)
key
=
strings
.
TrimPrefix
(
key
,
"sk-"
)
parts
:=
strings
.
Split
(
key
,
"-"
)
key
=
parts
[
0
]
}
token
,
err
:=
model
.
ValidateUserToken
(
key
)
token
,
err
:=
model
.
ValidateUserToken
(
key
)
if
err
!=
nil
{
if
err
!=
nil
{
c
.
JSON
(
http
.
StatusUnauthorized
,
gin
.
H
{
c
.
JSON
(
http
.
StatusUnauthorized
,
gin
.
H
{
...
...
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