Skip to content

Commit 8af4e28

Browse files
committed
feat: support cohere rerank
1 parent afe02c6 commit 8af4e28

25 files changed

Lines changed: 347 additions & 11 deletions

File tree

controller/relay.go

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -29,6 +29,8 @@ func relayHandler(c *gin.Context, relayMode int) *dto.OpenAIErrorWithStatusCode
2929
fallthrough
3030
case relayconstant.RelayModeAudioTranscription:
3131
err = relay.AudioHelper(c, relayMode)
32+
case relayconstant.RelayModeRerank:
33+
err = relay.RerankHelper(c, relayMode)
3234
default:
3335
err = relay.TextHelper(c)
3436
}

dto/rerank.go

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,19 @@
1+
package dto
2+
3+
type RerankRequest struct {
4+
Documents []any `json:"documents"`
5+
Query string `json:"query"`
6+
Model string `json:"model"`
7+
TopN int `json:"top_n"`
8+
}
9+
10+
type RerankResponseDocument struct {
11+
Document any `json:"document"`
12+
Index int `json:"index"`
13+
RelevanceScore float64 `json:"relevance_score"`
14+
}
15+
16+
type RerankResponse struct {
17+
Results []RerankResponseDocument `json:"results"`
18+
Usage Usage `json:"usage"`
19+
}

relay/channel/adapter.go

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -11,9 +11,11 @@ import (
1111
type Adaptor interface {
1212
// Init IsStream bool
1313
Init(info *relaycommon.RelayInfo, request dto.GeneralOpenAIRequest)
14+
InitRerank(info *relaycommon.RelayInfo, request dto.RerankRequest)
1415
GetRequestURL(info *relaycommon.RelayInfo) (string, error)
1516
SetupRequestHeader(c *gin.Context, req *http.Request, info *relaycommon.RelayInfo) error
1617
ConvertRequest(c *gin.Context, relayMode int, request *dto.GeneralOpenAIRequest) (any, error)
18+
ConvertRerankRequest(c *gin.Context, relayMode int, request dto.RerankRequest) (any, error)
1719
DoRequest(c *gin.Context, info *relaycommon.RelayInfo, requestBody io.Reader) (*http.Response, error)
1820
DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage *dto.Usage, err *dto.OpenAIErrorWithStatusCode)
1921
GetModelList() []string

relay/channel/ali/adaptor.go

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -15,6 +15,9 @@ import (
1515
type Adaptor struct {
1616
}
1717

18+
func (a *Adaptor) InitRerank(info *relaycommon.RelayInfo, request dto.RerankRequest) {
19+
}
20+
1821
func (a *Adaptor) Init(info *relaycommon.RelayInfo, request dto.GeneralOpenAIRequest) {
1922

2023
}
@@ -53,6 +56,10 @@ func (a *Adaptor) ConvertRequest(c *gin.Context, relayMode int, request *dto.Gen
5356
}
5457
}
5558

59+
func (a *Adaptor) ConvertRerankRequest(c *gin.Context, relayMode int, request dto.RerankRequest) (any, error) {
60+
return nil, nil
61+
}
62+
5663
func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, requestBody io.Reader) (*http.Response, error) {
5764
return channel.DoApiRequest(a, c, info, requestBody)
5865
}

relay/channel/aws/adaptor.go

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,11 @@ type Adaptor struct {
2020
RequestMode int
2121
}
2222

23+
func (a *Adaptor) InitRerank(info *relaycommon.RelayInfo, request dto.RerankRequest) {
24+
//TODO implement me
25+
26+
}
27+
2328
func (a *Adaptor) Init(info *relaycommon.RelayInfo, request dto.GeneralOpenAIRequest) {
2429
if strings.HasPrefix(info.UpstreamModelName, "claude-3") {
2530
a.RequestMode = RequestModeMessage
@@ -53,6 +58,10 @@ func (a *Adaptor) ConvertRequest(c *gin.Context, relayMode int, request *dto.Gen
5358
return claudeReq, err
5459
}
5560

61+
func (a *Adaptor) ConvertRerankRequest(c *gin.Context, relayMode int, request dto.RerankRequest) (any, error) {
62+
return nil, nil
63+
}
64+
5665
func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, requestBody io.Reader) (*http.Response, error) {
5766
return nil, nil
5867
}

relay/channel/baidu/adaptor.go

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,11 @@ import (
1616
type Adaptor struct {
1717
}
1818

19+
func (a *Adaptor) InitRerank(info *relaycommon.RelayInfo, request dto.RerankRequest) {
20+
//TODO implement me
21+
22+
}
23+
1924
func (a *Adaptor) Init(info *relaycommon.RelayInfo, request dto.GeneralOpenAIRequest) {
2025

2126
}
@@ -108,6 +113,10 @@ func (a *Adaptor) ConvertRequest(c *gin.Context, relayMode int, request *dto.Gen
108113
}
109114
}
110115

116+
func (a *Adaptor) ConvertRerankRequest(c *gin.Context, relayMode int, request dto.RerankRequest) (any, error) {
117+
return nil, nil
118+
}
119+
111120
func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, requestBody io.Reader) (*http.Response, error) {
112121
return channel.DoApiRequest(a, c, info, requestBody)
113122
}

relay/channel/claude/adaptor.go

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,11 @@ type Adaptor struct {
2121
RequestMode int
2222
}
2323

24+
func (a *Adaptor) InitRerank(info *relaycommon.RelayInfo, request dto.RerankRequest) {
25+
//TODO implement me
26+
27+
}
28+
2429
func (a *Adaptor) Init(info *relaycommon.RelayInfo, request dto.GeneralOpenAIRequest) {
2530
if strings.HasPrefix(info.UpstreamModelName, "claude-3") {
2631
a.RequestMode = RequestModeMessage
@@ -59,6 +64,10 @@ func (a *Adaptor) ConvertRequest(c *gin.Context, relayMode int, request *dto.Gen
5964
}
6065
}
6166

67+
func (a *Adaptor) ConvertRerankRequest(c *gin.Context, relayMode int, request dto.RerankRequest) (any, error) {
68+
return nil, nil
69+
}
70+
6271
func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, requestBody io.Reader) (*http.Response, error) {
6372
return channel.DoApiRequest(a, c, info, requestBody)
6473
}

relay/channel/cohere/adaptor.go

Lines changed: 20 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -8,16 +8,24 @@ import (
88
"one-api/dto"
99
"one-api/relay/channel"
1010
relaycommon "one-api/relay/common"
11+
"one-api/relay/constant"
1112
)
1213

1314
type Adaptor struct {
1415
}
1516

17+
func (a *Adaptor) InitRerank(info *relaycommon.RelayInfo, request dto.RerankRequest) {
18+
}
19+
1620
func (a *Adaptor) Init(info *relaycommon.RelayInfo, request dto.GeneralOpenAIRequest) {
1721
}
1822

1923
func (a *Adaptor) GetRequestURL(info *relaycommon.RelayInfo) (string, error) {
20-
return fmt.Sprintf("%s/v1/chat", info.BaseUrl), nil
24+
if info.RelayMode == constant.RelayModeRerank {
25+
return fmt.Sprintf("%s/v1/rerank", info.BaseUrl), nil
26+
} else {
27+
return fmt.Sprintf("%s/v1/chat", info.BaseUrl), nil
28+
}
2129
}
2230

2331
func (a *Adaptor) SetupRequestHeader(c *gin.Context, req *http.Request, info *relaycommon.RelayInfo) error {
@@ -34,11 +42,19 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request
3442
return channel.DoApiRequest(a, c, info, requestBody)
3543
}
3644

45+
func (a *Adaptor) ConvertRerankRequest(c *gin.Context, relayMode int, request dto.RerankRequest) (any, error) {
46+
return requestConvertRerank2Cohere(request), nil
47+
}
48+
3749
func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage *dto.Usage, err *dto.OpenAIErrorWithStatusCode) {
38-
if info.IsStream {
39-
err, usage = cohereStreamHandler(c, resp, info)
50+
if info.RelayMode == constant.RelayModeRerank {
51+
err, usage = cohereRerankHandler(c, resp, info)
4052
} else {
41-
err, usage = cohereHandler(c, resp, info.UpstreamModelName, info.PromptTokens)
53+
if info.IsStream {
54+
err, usage = cohereStreamHandler(c, resp, info)
55+
} else {
56+
err, usage = cohereHandler(c, resp, info.UpstreamModelName, info.PromptTokens)
57+
}
4258
}
4359
return
4460
}

relay/channel/cohere/constant.go

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@ package cohere
22

33
var ModelList = []string{
44
"command-r", "command-r-plus", "command-light", "command-light-nightly", "command", "command-nightly",
5+
"rerank-english-v3.0", "rerank-multilingual-v3.0", "rerank-english-v2.0", "rerank-multilingual-v2.0",
56
}
67

78
var ChannelName = "cohere"

relay/channel/cohere/dto.go

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,7 @@
11
package cohere
22

3+
import "one-api/dto"
4+
35
type CohereRequest struct {
46
Model string `json:"model"`
57
ChatHistory []ChatHistory `json:"chat_history"`
@@ -28,6 +30,19 @@ type CohereResponseResult struct {
2830
Meta CohereMeta `json:"meta"`
2931
}
3032

33+
type CohereRerankRequest struct {
34+
Documents []any `json:"documents"`
35+
Query string `json:"query"`
36+
Model string `json:"model"`
37+
TopN int `json:"top_n"`
38+
ReturnDocuments bool `json:"return_documents"`
39+
}
40+
41+
type CohereRerankResponseResult struct {
42+
Results []dto.RerankResponseDocument `json:"results"`
43+
Meta CohereMeta `json:"meta"`
44+
}
45+
3146
type CohereMeta struct {
3247
//Tokens CohereTokens `json:"tokens"`
3348
BilledUnits CohereBilledUnits `json:"billed_units"`

0 commit comments

Comments
 (0)