diff --git a/waveshell/pkg/packet/packet.go b/waveshell/pkg/packet/packet.go index 435c464b..c2dec050 100644 --- a/waveshell/pkg/packet/packet.go +++ b/waveshell/pkg/packet/packet.go @@ -60,7 +60,8 @@ const ( WriteFileDonePacketStr = "writefiledone" // rpc-response FileDataPacketStr = "filedata" - OpenAIPacketStr = "openai" // other + OpenAIPacketStr = "openai" // other + OpenAICloudReqStr = "openai-cloudreq" ) const PacketSenderQueueSize = 20 @@ -842,6 +843,30 @@ func MakeWriteFileDonePacket(reqId string) *WriteFileDonePacketType { } } +type OpenAIPromptMessageType struct { + Role string `json:"role"` + Content string `json:"content"` + Name string `json:"name,omitempty"` +} + +type OpenAICloudReqPacketType struct { + Type string `json:"type"` + ClientId string `json:"clientid"` + Prompt []OpenAIPromptMessageType `json:"prompt"` + MaxTokens int `json:"maxtokens,omitempty"` + MaxChoices int `json:"maxchoices,omitempty"` +} + +func (*OpenAICloudReqPacketType) GetType() string { + return OpenAICloudReqStr +} + +func MakeOpenAICloudReqPacket() *OpenAICloudReqPacketType { + return &OpenAICloudReqPacketType{ + Type: OpenAICloudReqStr, + } +} + type PacketType interface { GetType() string } diff --git a/wavesrv/pkg/cmdrunner/cmdrunner.go b/wavesrv/pkg/cmdrunner/cmdrunner.go index 5bf581d2..d1bd79c4 100644 --- a/wavesrv/pkg/cmdrunner/cmdrunner.go +++ b/wavesrv/pkg/cmdrunner/cmdrunner.go @@ -1462,7 +1462,7 @@ func writePacketToPty(ctx context.Context, cmd *sstore.CmdType, pk packet.Packet return nil } -func doOpenAICompletion(cmd *sstore.CmdType, opts *sstore.OpenAIOptsType, prompt []sstore.OpenAIPromptMessageType) { +func doOpenAICompletion(cmd *sstore.CmdType, opts *sstore.OpenAIOptsType, prompt []packet.OpenAIPromptMessageType) { var outputPos int64 var hadError bool startTime := time.Now() @@ -1514,7 +1514,7 @@ func doOpenAICompletion(cmd *sstore.CmdType, opts *sstore.OpenAIOptsType, prompt return } -func doOpenAIStreamCompletion(cmd *sstore.CmdType, opts *sstore.OpenAIOptsType, prompt []sstore.OpenAIPromptMessageType) { +func doOpenAIStreamCompletion(cmd *sstore.CmdType, clientId string, opts *sstore.OpenAIOptsType, prompt []packet.OpenAIPromptMessageType) { var outputPos int64 var hadError bool startTime := time.Now() @@ -1552,7 +1552,7 @@ func doOpenAIStreamCompletion(cmd *sstore.CmdType, opts *sstore.OpenAIOptsType, var err error if opts.APIToken == "" { var conn *websocket.Conn - ch, conn, err = openai.RunCloudCompletionStream(ctx, opts, prompt) + ch, conn, err = openai.RunCloudCompletionStream(ctx, clientId, opts, prompt) if conn != nil { defer conn.Close() } @@ -1641,9 +1641,9 @@ func OpenAICommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstor if err != nil { return nil, fmt.Errorf("cannot add new line: %v", err) } - prompt := []sstore.OpenAIPromptMessageType{{Role: sstore.OpenAIRoleUser, Content: promptStr}} + prompt := []packet.OpenAIPromptMessageType{{Role: sstore.OpenAIRoleUser, Content: promptStr}} if resolveBool(pk.Kwargs["stream"], true) { - go doOpenAIStreamCompletion(cmd, opts, prompt) + go doOpenAIStreamCompletion(cmd, clientData.ClientId, opts, prompt) } else { go doOpenAICompletion(cmd, opts, prompt) } diff --git a/wavesrv/pkg/remote/openai/openai.go b/wavesrv/pkg/remote/openai/openai.go index 00d8ccc2..ed325172 100644 --- a/wavesrv/pkg/remote/openai/openai.go +++ b/wavesrv/pkg/remote/openai/openai.go @@ -37,7 +37,7 @@ func convertUsage(resp openaiapi.ChatCompletionResponse) *packet.OpenAIUsageType } } -func ConvertPrompt(prompt []sstore.OpenAIPromptMessageType) []openaiapi.ChatCompletionMessage { +func ConvertPrompt(prompt []packet.OpenAIPromptMessageType) []openaiapi.ChatCompletionMessage { var rtn []openaiapi.ChatCompletionMessage for _, p := range prompt { msg := openaiapi.ChatCompletionMessage{Role: p.Role, Content: p.Content, Name: p.Name} @@ -46,7 +46,7 @@ func ConvertPrompt(prompt []sstore.OpenAIPromptMessageType) []openaiapi.ChatComp return rtn } -func RunCompletion(ctx context.Context, opts *sstore.OpenAIOptsType, prompt []sstore.OpenAIPromptMessageType) ([]*packet.OpenAIPacketType, error) { +func RunCompletion(ctx context.Context, opts *sstore.OpenAIOptsType, prompt []packet.OpenAIPromptMessageType) ([]*packet.OpenAIPacketType, error) { if opts == nil { return nil, fmt.Errorf("no openai opts found") } @@ -79,21 +79,22 @@ func RunCompletion(ctx context.Context, opts *sstore.OpenAIOptsType, prompt []ss return marshalResponse(apiResp), nil } -func RunCloudCompletionStream(ctx context.Context, opts *sstore.OpenAIOptsType, prompt []sstore.OpenAIPromptMessageType) (chan *packet.OpenAIPacketType, *websocket.Conn, error) { +func RunCloudCompletionStream(ctx context.Context, clientId string, opts *sstore.OpenAIOptsType, prompt []packet.OpenAIPromptMessageType) (chan *packet.OpenAIPacketType, *websocket.Conn, error) { if opts == nil { return nil, nil, fmt.Errorf("no openai opts found") } - websocketContext, _ := context.WithTimeout(context.Background(), CloudWebsocketConnectTimeout) + websocketContext, dialCancelFn := context.WithTimeout(context.Background(), CloudWebsocketConnectTimeout) + defer dialCancelFn() conn, _, err := websocket.DefaultDialer.DialContext(websocketContext, pcloud.GetWSEndpoint(), nil) if err != nil { return nil, nil, fmt.Errorf("OpenAI request, websocket connect error: %v", err) } - cloudCompletionRequestConfig := sstore.OpenAICloudCompletionRequest{ - Prompt: prompt, - MaxTokens: opts.MaxTokens, - MaxChoices: opts.MaxChoices, - } - configMessageBuf, err := json.Marshal(cloudCompletionRequestConfig) + reqPk := packet.MakeOpenAICloudReqPacket() + reqPk.ClientId = clientId + reqPk.Prompt = prompt + reqPk.MaxTokens = opts.MaxTokens + reqPk.MaxChoices = opts.MaxChoices + configMessageBuf, err := json.Marshal(reqPk) err = conn.WriteMessage(websocket.TextMessage, configMessageBuf) if err != nil { return nil, nil, fmt.Errorf("OpenAI request, websocket write config error: %v", err) @@ -134,7 +135,7 @@ func RunCloudCompletionStream(ctx context.Context, opts *sstore.OpenAIOptsType, return rtn, conn, err } -func RunCompletionStream(ctx context.Context, opts *sstore.OpenAIOptsType, prompt []sstore.OpenAIPromptMessageType) (chan *packet.OpenAIPacketType, error) { +func RunCompletionStream(ctx context.Context, opts *sstore.OpenAIOptsType, prompt []packet.OpenAIPromptMessageType) (chan *packet.OpenAIPacketType, error) { if opts == nil { return nil, fmt.Errorf("no openai opts found") } diff --git a/wavesrv/pkg/sstore/sstore.go b/wavesrv/pkg/sstore/sstore.go index ff7f389b..527e06fc 100644 --- a/wavesrv/pkg/sstore/sstore.go +++ b/wavesrv/pkg/sstore/sstore.go @@ -778,18 +778,6 @@ type OpenAIResponse struct { Choices []OpenAIChoiceType `json:"choices,omitempty"` } -type OpenAIPromptMessageType struct { - Role string `json:"role"` - Content string `json:"content"` - Name string `json:"name,omitempty"` -} - -type OpenAICloudCompletionRequest struct { - Prompt []OpenAIPromptMessageType `json:"prompt"` - MaxTokens int `json:"maxtokens,omitempty"` - MaxChoices int `json:"maxchoices,omitempty"` -} - type PlaybookType struct { PlaybookId string `json:"playbookid"` PlaybookName string `json:"playbookname"`