| 99 | } |
| 100 | |
| 101 | func (p *OpenAI) CreateInterceptor(_ http.ResponseWriter, r *http.Request, tracer trace.Tracer) (_ intercept.Interceptor, outErr error) { |
| 102 | id := uuid.New() |
| 103 | |
| 104 | _, span := tracer.Start(r.Context(), "Intercept.CreateInterceptor") |
| 105 | defer tracing.EndSpanErr(span, &outErr) |
| 106 | |
| 107 | var interceptor intercept.Interceptor |
| 108 | |
| 109 | cfg := p.cfg |
| 110 | // At this point the request contains only LLM provider headers. Any |
| 111 | // Coder-specific authentication has already been stripped. |
| 112 | // |
| 113 | // In centralized mode Authorization is absent, so cfg keeps the |
| 114 | // centralized key unchanged. |
| 115 | // |
| 116 | // In BYOK mode the user's credential is in Authorization. Replace |
| 117 | // the centralized key with it so it is forwarded upstream. |
| 118 | credKind := intercept.CredentialKindCentralized |
| 119 | if token := utils.ExtractBearerToken(r.Header.Get("Authorization")); token != "" { |
| 120 | cfg.Key = token |
| 121 | credKind = intercept.CredentialKindBYOK |
| 122 | } |
| 123 | cred := intercept.NewCredentialInfo(credKind, cfg.Key) |
| 124 | |
| 125 | path := strings.TrimPrefix(r.URL.Path, p.RoutePrefix()) |
| 126 | switch path { |
| 127 | case routeChatCompletions: |
| 128 | var req chatcompletions.ChatCompletionNewParamsWrapper |
| 129 | if err := json.NewDecoder(r.Body).Decode(&req); err != nil { |
| 130 | return nil, xerrors.Errorf("unmarshal request body: %w", err) |
| 131 | } |
| 132 | |
| 133 | if req.Stream { |
| 134 | interceptor = chatcompletions.NewStreamingInterceptor(id, &req, p.Name(), cfg, r.Header, p.AuthHeader(), tracer, cred) |
| 135 | } else { |
| 136 | interceptor = chatcompletions.NewBlockingInterceptor(id, &req, p.Name(), cfg, r.Header, p.AuthHeader(), tracer, cred) |
| 137 | } |
| 138 | |
| 139 | case routeResponses: |
| 140 | payload, err := io.ReadAll(r.Body) |
| 141 | if err != nil { |
| 142 | return nil, xerrors.Errorf("read body: %w", err) |
| 143 | } |
| 144 | reqPayload, err := responses.NewRequestPayload(payload) |
| 145 | if err != nil { |
| 146 | return nil, xerrors.Errorf("unmarshal request body: %w", err) |
| 147 | } |
| 148 | if reqPayload.Stream() { |
| 149 | interceptor = responses.NewStreamingInterceptor(id, reqPayload, p.Name(), cfg, r.Header, p.AuthHeader(), tracer, cred) |
| 150 | } else { |
| 151 | interceptor = responses.NewBlockingInterceptor(id, reqPayload, p.Name(), cfg, r.Header, p.AuthHeader(), tracer, cred) |
| 152 | } |
| 153 | |
| 154 | default: |
| 155 | span.SetStatus(codes.Error, "unknown route: "+r.URL.Path) |
| 156 | return nil, ErrUnknownRoute |
| 157 | } |
| 158 | span.SetAttributes(interceptor.TraceAttributes(r)...) |