128 lines
4.1 KiB
Go
128 lines
4.1 KiB
Go
package openai
|
|
|
|
import (
|
|
"context"
|
|
"fmt"
|
|
"net/http"
|
|
"sync"
|
|
|
|
edgeservice "iop/apps/edge/internal/service"
|
|
iop "iop/proto/gen/iop"
|
|
)
|
|
|
|
// newTestChatDispatchContext builds the minimal dispatch context a chat
|
|
// completion stage needs when a test drives that stage directly instead of
|
|
// going through the ingress handler.
|
|
func newTestChatDispatchContext(r *http.Request, req chatCompletionRequest) *chatDispatchContext {
|
|
return &chatDispatchContext{
|
|
openAIRequestContext: openAIRequestContext{
|
|
r: r,
|
|
endpoint: usageEndpointChatCompletions,
|
|
},
|
|
req: req,
|
|
retrySubmit: func(ctx context.Context, retryReq edgeservice.SubmitRunRequest) (any, error) {
|
|
return nil, nil
|
|
},
|
|
}
|
|
}
|
|
|
|
// fakeRunService is the small default double for the normalized run, cancel,
|
|
// and Ollama surfaces. Provider tunnel/pool scenarios use
|
|
// providerFakeRunService, which embeds this type.
|
|
type fakeRunService struct {
|
|
req edgeservice.SubmitRunRequest
|
|
reqs []edgeservice.SubmitRunRequest
|
|
ollamaReq edgeservice.OllamaAPIRequest
|
|
ollamaResp edgeservice.OllamaAPIView
|
|
events chan *iop.RunEvent
|
|
eventRuns []chan *iop.RunEvent
|
|
runIDs []string
|
|
|
|
submitMu sync.Mutex
|
|
submitErrAfter int // zero means no error; >=1 means error starts from this req index (1-based)
|
|
cancelMu sync.Mutex
|
|
cancelCalls []edgeservice.CancelRunRequest
|
|
}
|
|
|
|
// SubmitProviderTunnel fails fast: a fixture reaching the provider tunnel path
|
|
// must use providerFakeRunService.
|
|
func (s *fakeRunService) SubmitProviderTunnel(context.Context, edgeservice.SubmitProviderTunnelRequest) (edgeservice.ProviderTunnelResult, error) {
|
|
return nil, fmt.Errorf("fakeRunService: provider tunnel dispatch is not configured; use providerFakeRunService")
|
|
}
|
|
|
|
// SubmitProviderPool fails fast: a fixture reaching the provider-pool path must
|
|
// use providerFakeRunService.
|
|
func (s *fakeRunService) SubmitProviderPool(context.Context, edgeservice.ProviderPoolDispatchRequest) (*edgeservice.ProviderPoolDispatchResult, error) {
|
|
return nil, fmt.Errorf("fakeRunService: provider pool dispatch is not configured; use providerFakeRunService")
|
|
}
|
|
|
|
func (s *fakeRunService) reqsSnapshot() []edgeservice.SubmitRunRequest {
|
|
s.submitMu.Lock()
|
|
defer s.submitMu.Unlock()
|
|
return append([]edgeservice.SubmitRunRequest(nil), s.reqs...)
|
|
}
|
|
|
|
func (s *fakeRunService) CancelRun(_ context.Context, req edgeservice.CancelRunRequest) (edgeservice.CommandResult, error) {
|
|
s.cancelMu.Lock()
|
|
s.cancelCalls = append(s.cancelCalls, req)
|
|
s.cancelMu.Unlock()
|
|
return edgeservice.CommandResult{NodeID: req.NodeRef, SessionID: req.SessionID}, nil
|
|
}
|
|
|
|
func (s *fakeRunService) cancelCallsSnapshot() []edgeservice.CancelRunRequest {
|
|
s.cancelMu.Lock()
|
|
defer s.cancelMu.Unlock()
|
|
return append([]edgeservice.CancelRunRequest(nil), s.cancelCalls...)
|
|
}
|
|
|
|
func (s *fakeRunService) SubmitRun(_ context.Context, req edgeservice.SubmitRunRequest) (edgeservice.RunResult, error) {
|
|
s.submitMu.Lock()
|
|
s.req = req
|
|
s.reqs = append(s.reqs, req)
|
|
attempt := len(s.reqs)
|
|
if s.submitErrAfter > 0 && attempt >= s.submitErrAfter {
|
|
s.submitMu.Unlock()
|
|
return nil, fmt.Errorf("inject error after %d attempts", attempt-1)
|
|
}
|
|
events := s.events
|
|
if attempt-1 < len(s.eventRuns) {
|
|
events = s.eventRuns[attempt-1]
|
|
}
|
|
runID := "run-test"
|
|
if attempt-1 < len(s.runIDs) {
|
|
runID = s.runIDs[attempt-1]
|
|
} else if attempt > 1 {
|
|
runID = fmt.Sprintf("run-test-%d", attempt)
|
|
}
|
|
s.submitMu.Unlock()
|
|
return &edgeservice.RunHandle{
|
|
RunDispatch: edgeservice.RunDispatch{
|
|
RunID: runID,
|
|
ModelGroupKey: req.ModelGroupKey,
|
|
Adapter: req.Adapter,
|
|
Target: req.Target,
|
|
SessionID: req.SessionID,
|
|
TimeoutSec: 5,
|
|
},
|
|
RunStream: edgeservice.RunStream{
|
|
Events: events,
|
|
NodeEvents: make(chan *iop.EdgeNodeEvent),
|
|
},
|
|
}, nil
|
|
}
|
|
|
|
func (s *fakeRunService) OllamaAPI(_ context.Context, req edgeservice.OllamaAPIRequest) (edgeservice.OllamaAPIView, error) {
|
|
s.ollamaReq = req
|
|
if s.ollamaResp.StatusCode == 0 {
|
|
s.ollamaResp.StatusCode = http.StatusOK
|
|
}
|
|
return s.ollamaResp, nil
|
|
}
|
|
|
|
func bufferedRunEvents(events ...*iop.RunEvent) chan *iop.RunEvent {
|
|
ch := make(chan *iop.RunEvent, len(events))
|
|
for _, event := range events {
|
|
ch <- event
|
|
}
|
|
return ch
|
|
}
|