mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: add aibridgedserver pkg (#19902)
This commit is contained in:
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,114 @@
|
||||
syntax = "proto3";
|
||||
option go_package = "github.com/coder/coder/v2/aibridged/proto";
|
||||
|
||||
package proto;
|
||||
|
||||
import "google/protobuf/any.proto";
|
||||
import "google/protobuf/timestamp.proto";
|
||||
|
||||
// Recorder is responsible for persisting AI usage records along with their related interception.
|
||||
service Recorder {
|
||||
// RecordInterception creates a new interception record to which all other sub-resources
|
||||
// (token, prompt, tool uses) will be related.
|
||||
rpc RecordInterception(RecordInterceptionRequest) returns (RecordInterceptionResponse);
|
||||
rpc RecordTokenUsage(RecordTokenUsageRequest) returns (RecordTokenUsageResponse);
|
||||
rpc RecordPromptUsage(RecordPromptUsageRequest) returns (RecordPromptUsageResponse);
|
||||
rpc RecordToolUsage(RecordToolUsageRequest) returns (RecordToolUsageResponse);
|
||||
}
|
||||
|
||||
// MCPConfigurator is responsible for retrieving any relevant data required for configuring MCP clients
|
||||
// against remote servers.
|
||||
service MCPConfigurator {
|
||||
// GetMCPServerConfigs will retrieve MCP server configurations.
|
||||
rpc GetMCPServerConfigs(GetMCPServerConfigsRequest) returns (GetMCPServerConfigsResponse);
|
||||
// GetMCPServerAccessTokensBatch will retrieve an access token for a given list of MCP servers, which may involve
|
||||
// acquiring, validating, or refreshing tokens synchronously. The server should make every effort to
|
||||
// parallelise this work.
|
||||
rpc GetMCPServerAccessTokensBatch(GetMCPServerAccessTokensBatchRequest) returns (GetMCPServerAccessTokensBatchResponse);
|
||||
}
|
||||
|
||||
// Authorizer handles all Coder-related authorization functions.
|
||||
service Authorizer {
|
||||
// IsAuthorized validates that a given Coder key is valid and the user is authorized to use AI Bridge.
|
||||
// TODO: add authorization; currently only key validation takes place.
|
||||
rpc IsAuthorized(IsAuthorizedRequest) returns (IsAuthorizedResponse);
|
||||
}
|
||||
|
||||
message RecordInterceptionRequest {
|
||||
string id = 1; // UUID.
|
||||
string initiator_id = 2; // UUID.
|
||||
string provider = 3;
|
||||
string model = 4;
|
||||
map<string, google.protobuf.Any> metadata = 5;
|
||||
google.protobuf.Timestamp started_at = 6;
|
||||
}
|
||||
|
||||
message RecordInterceptionResponse {}
|
||||
|
||||
message RecordTokenUsageRequest {
|
||||
string interception_id = 1; // UUID.
|
||||
string msg_id = 2; // ID provided by provider.
|
||||
int64 input_tokens = 3;
|
||||
int64 output_tokens = 4;
|
||||
map<string, google.protobuf.Any> metadata = 5;
|
||||
google.protobuf.Timestamp created_at = 6;
|
||||
}
|
||||
message RecordTokenUsageResponse {}
|
||||
|
||||
message RecordPromptUsageRequest {
|
||||
string interception_id = 1; // UUID.
|
||||
string msg_id = 2; // ID provided by provider.
|
||||
string prompt = 3;
|
||||
map<string, google.protobuf.Any> metadata = 4;
|
||||
google.protobuf.Timestamp created_at = 5;
|
||||
}
|
||||
message RecordPromptUsageResponse {}
|
||||
|
||||
message RecordToolUsageRequest {
|
||||
string interception_id = 1; // UUID.
|
||||
string msg_id = 2; // ID provided by provider.
|
||||
optional string server_url = 3; // The URL of the MCP server.
|
||||
string tool = 4;
|
||||
string input = 5;
|
||||
bool injected = 6;
|
||||
optional string invocation_error = 7; // Only injected tools are invoked.
|
||||
map<string, google.protobuf.Any> metadata = 8;
|
||||
google.protobuf.Timestamp created_at = 9;
|
||||
}
|
||||
message RecordToolUsageResponse {}
|
||||
|
||||
message GetMCPServerConfigsRequest {
|
||||
string user_id = 1; // UUID. // Not used yet, will be necessary for later RBAC purposes.
|
||||
}
|
||||
|
||||
message GetMCPServerConfigsResponse {
|
||||
MCPServerConfig coder_mcp_config = 1;
|
||||
repeated MCPServerConfig external_auth_mcp_configs = 2;
|
||||
}
|
||||
|
||||
message MCPServerConfig {
|
||||
string id = 1; // Maps to the ID of the External Auth; this ID is unique.
|
||||
string url = 2;
|
||||
string tool_allow_regex = 3;
|
||||
string tool_deny_regex = 4;
|
||||
}
|
||||
|
||||
message GetMCPServerAccessTokensBatchRequest {
|
||||
string user_id = 1; // UUID.
|
||||
repeated string mcp_server_config_ids = 2;
|
||||
}
|
||||
|
||||
// GetMCPServerAccessTokensBatchResponse returns a map for resulting tokens or errors, indexed
|
||||
// by server ID.
|
||||
message GetMCPServerAccessTokensBatchResponse{
|
||||
map<string, string> access_tokens = 1;
|
||||
map<string, string> errors = 2;
|
||||
}
|
||||
|
||||
message IsAuthorizedRequest {
|
||||
string key = 1;
|
||||
}
|
||||
|
||||
message IsAuthorizedResponse {
|
||||
string owner_id = 1;
|
||||
}
|
||||
@@ -0,0 +1,421 @@
|
||||
// Code generated by protoc-gen-go-drpc. DO NOT EDIT.
|
||||
// protoc-gen-go-drpc version: v0.0.34
|
||||
// source: enterprise/x/aibridged/proto/aibridged.proto
|
||||
|
||||
package proto
|
||||
|
||||
import (
|
||||
context "context"
|
||||
errors "errors"
|
||||
protojson "google.golang.org/protobuf/encoding/protojson"
|
||||
proto "google.golang.org/protobuf/proto"
|
||||
drpc "storj.io/drpc"
|
||||
drpcerr "storj.io/drpc/drpcerr"
|
||||
)
|
||||
|
||||
type drpcEncoding_File_enterprise_x_aibridged_proto_aibridged_proto struct{}
|
||||
|
||||
func (drpcEncoding_File_enterprise_x_aibridged_proto_aibridged_proto) Marshal(msg drpc.Message) ([]byte, error) {
|
||||
return proto.Marshal(msg.(proto.Message))
|
||||
}
|
||||
|
||||
func (drpcEncoding_File_enterprise_x_aibridged_proto_aibridged_proto) MarshalAppend(buf []byte, msg drpc.Message) ([]byte, error) {
|
||||
return proto.MarshalOptions{}.MarshalAppend(buf, msg.(proto.Message))
|
||||
}
|
||||
|
||||
func (drpcEncoding_File_enterprise_x_aibridged_proto_aibridged_proto) Unmarshal(buf []byte, msg drpc.Message) error {
|
||||
return proto.Unmarshal(buf, msg.(proto.Message))
|
||||
}
|
||||
|
||||
func (drpcEncoding_File_enterprise_x_aibridged_proto_aibridged_proto) JSONMarshal(msg drpc.Message) ([]byte, error) {
|
||||
return protojson.Marshal(msg.(proto.Message))
|
||||
}
|
||||
|
||||
func (drpcEncoding_File_enterprise_x_aibridged_proto_aibridged_proto) JSONUnmarshal(buf []byte, msg drpc.Message) error {
|
||||
return protojson.Unmarshal(buf, msg.(proto.Message))
|
||||
}
|
||||
|
||||
type DRPCRecorderClient interface {
|
||||
DRPCConn() drpc.Conn
|
||||
|
||||
RecordInterception(ctx context.Context, in *RecordInterceptionRequest) (*RecordInterceptionResponse, error)
|
||||
RecordTokenUsage(ctx context.Context, in *RecordTokenUsageRequest) (*RecordTokenUsageResponse, error)
|
||||
RecordPromptUsage(ctx context.Context, in *RecordPromptUsageRequest) (*RecordPromptUsageResponse, error)
|
||||
RecordToolUsage(ctx context.Context, in *RecordToolUsageRequest) (*RecordToolUsageResponse, error)
|
||||
}
|
||||
|
||||
type drpcRecorderClient struct {
|
||||
cc drpc.Conn
|
||||
}
|
||||
|
||||
func NewDRPCRecorderClient(cc drpc.Conn) DRPCRecorderClient {
|
||||
return &drpcRecorderClient{cc}
|
||||
}
|
||||
|
||||
func (c *drpcRecorderClient) DRPCConn() drpc.Conn { return c.cc }
|
||||
|
||||
func (c *drpcRecorderClient) RecordInterception(ctx context.Context, in *RecordInterceptionRequest) (*RecordInterceptionResponse, error) {
|
||||
out := new(RecordInterceptionResponse)
|
||||
err := c.cc.Invoke(ctx, "/proto.Recorder/RecordInterception", drpcEncoding_File_enterprise_x_aibridged_proto_aibridged_proto{}, in, out)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (c *drpcRecorderClient) RecordTokenUsage(ctx context.Context, in *RecordTokenUsageRequest) (*RecordTokenUsageResponse, error) {
|
||||
out := new(RecordTokenUsageResponse)
|
||||
err := c.cc.Invoke(ctx, "/proto.Recorder/RecordTokenUsage", drpcEncoding_File_enterprise_x_aibridged_proto_aibridged_proto{}, in, out)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (c *drpcRecorderClient) RecordPromptUsage(ctx context.Context, in *RecordPromptUsageRequest) (*RecordPromptUsageResponse, error) {
|
||||
out := new(RecordPromptUsageResponse)
|
||||
err := c.cc.Invoke(ctx, "/proto.Recorder/RecordPromptUsage", drpcEncoding_File_enterprise_x_aibridged_proto_aibridged_proto{}, in, out)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (c *drpcRecorderClient) RecordToolUsage(ctx context.Context, in *RecordToolUsageRequest) (*RecordToolUsageResponse, error) {
|
||||
out := new(RecordToolUsageResponse)
|
||||
err := c.cc.Invoke(ctx, "/proto.Recorder/RecordToolUsage", drpcEncoding_File_enterprise_x_aibridged_proto_aibridged_proto{}, in, out)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
type DRPCRecorderServer interface {
|
||||
RecordInterception(context.Context, *RecordInterceptionRequest) (*RecordInterceptionResponse, error)
|
||||
RecordTokenUsage(context.Context, *RecordTokenUsageRequest) (*RecordTokenUsageResponse, error)
|
||||
RecordPromptUsage(context.Context, *RecordPromptUsageRequest) (*RecordPromptUsageResponse, error)
|
||||
RecordToolUsage(context.Context, *RecordToolUsageRequest) (*RecordToolUsageResponse, error)
|
||||
}
|
||||
|
||||
type DRPCRecorderUnimplementedServer struct{}
|
||||
|
||||
func (s *DRPCRecorderUnimplementedServer) RecordInterception(context.Context, *RecordInterceptionRequest) (*RecordInterceptionResponse, error) {
|
||||
return nil, drpcerr.WithCode(errors.New("Unimplemented"), drpcerr.Unimplemented)
|
||||
}
|
||||
|
||||
func (s *DRPCRecorderUnimplementedServer) RecordTokenUsage(context.Context, *RecordTokenUsageRequest) (*RecordTokenUsageResponse, error) {
|
||||
return nil, drpcerr.WithCode(errors.New("Unimplemented"), drpcerr.Unimplemented)
|
||||
}
|
||||
|
||||
func (s *DRPCRecorderUnimplementedServer) RecordPromptUsage(context.Context, *RecordPromptUsageRequest) (*RecordPromptUsageResponse, error) {
|
||||
return nil, drpcerr.WithCode(errors.New("Unimplemented"), drpcerr.Unimplemented)
|
||||
}
|
||||
|
||||
func (s *DRPCRecorderUnimplementedServer) RecordToolUsage(context.Context, *RecordToolUsageRequest) (*RecordToolUsageResponse, error) {
|
||||
return nil, drpcerr.WithCode(errors.New("Unimplemented"), drpcerr.Unimplemented)
|
||||
}
|
||||
|
||||
type DRPCRecorderDescription struct{}
|
||||
|
||||
func (DRPCRecorderDescription) NumMethods() int { return 4 }
|
||||
|
||||
func (DRPCRecorderDescription) Method(n int) (string, drpc.Encoding, drpc.Receiver, interface{}, bool) {
|
||||
switch n {
|
||||
case 0:
|
||||
return "/proto.Recorder/RecordInterception", drpcEncoding_File_enterprise_x_aibridged_proto_aibridged_proto{},
|
||||
func(srv interface{}, ctx context.Context, in1, in2 interface{}) (drpc.Message, error) {
|
||||
return srv.(DRPCRecorderServer).
|
||||
RecordInterception(
|
||||
ctx,
|
||||
in1.(*RecordInterceptionRequest),
|
||||
)
|
||||
}, DRPCRecorderServer.RecordInterception, true
|
||||
case 1:
|
||||
return "/proto.Recorder/RecordTokenUsage", drpcEncoding_File_enterprise_x_aibridged_proto_aibridged_proto{},
|
||||
func(srv interface{}, ctx context.Context, in1, in2 interface{}) (drpc.Message, error) {
|
||||
return srv.(DRPCRecorderServer).
|
||||
RecordTokenUsage(
|
||||
ctx,
|
||||
in1.(*RecordTokenUsageRequest),
|
||||
)
|
||||
}, DRPCRecorderServer.RecordTokenUsage, true
|
||||
case 2:
|
||||
return "/proto.Recorder/RecordPromptUsage", drpcEncoding_File_enterprise_x_aibridged_proto_aibridged_proto{},
|
||||
func(srv interface{}, ctx context.Context, in1, in2 interface{}) (drpc.Message, error) {
|
||||
return srv.(DRPCRecorderServer).
|
||||
RecordPromptUsage(
|
||||
ctx,
|
||||
in1.(*RecordPromptUsageRequest),
|
||||
)
|
||||
}, DRPCRecorderServer.RecordPromptUsage, true
|
||||
case 3:
|
||||
return "/proto.Recorder/RecordToolUsage", drpcEncoding_File_enterprise_x_aibridged_proto_aibridged_proto{},
|
||||
func(srv interface{}, ctx context.Context, in1, in2 interface{}) (drpc.Message, error) {
|
||||
return srv.(DRPCRecorderServer).
|
||||
RecordToolUsage(
|
||||
ctx,
|
||||
in1.(*RecordToolUsageRequest),
|
||||
)
|
||||
}, DRPCRecorderServer.RecordToolUsage, true
|
||||
default:
|
||||
return "", nil, nil, nil, false
|
||||
}
|
||||
}
|
||||
|
||||
func DRPCRegisterRecorder(mux drpc.Mux, impl DRPCRecorderServer) error {
|
||||
return mux.Register(impl, DRPCRecorderDescription{})
|
||||
}
|
||||
|
||||
type DRPCRecorder_RecordInterceptionStream interface {
|
||||
drpc.Stream
|
||||
SendAndClose(*RecordInterceptionResponse) error
|
||||
}
|
||||
|
||||
type drpcRecorder_RecordInterceptionStream struct {
|
||||
drpc.Stream
|
||||
}
|
||||
|
||||
func (x *drpcRecorder_RecordInterceptionStream) SendAndClose(m *RecordInterceptionResponse) error {
|
||||
if err := x.MsgSend(m, drpcEncoding_File_enterprise_x_aibridged_proto_aibridged_proto{}); err != nil {
|
||||
return err
|
||||
}
|
||||
return x.CloseSend()
|
||||
}
|
||||
|
||||
type DRPCRecorder_RecordTokenUsageStream interface {
|
||||
drpc.Stream
|
||||
SendAndClose(*RecordTokenUsageResponse) error
|
||||
}
|
||||
|
||||
type drpcRecorder_RecordTokenUsageStream struct {
|
||||
drpc.Stream
|
||||
}
|
||||
|
||||
func (x *drpcRecorder_RecordTokenUsageStream) SendAndClose(m *RecordTokenUsageResponse) error {
|
||||
if err := x.MsgSend(m, drpcEncoding_File_enterprise_x_aibridged_proto_aibridged_proto{}); err != nil {
|
||||
return err
|
||||
}
|
||||
return x.CloseSend()
|
||||
}
|
||||
|
||||
type DRPCRecorder_RecordPromptUsageStream interface {
|
||||
drpc.Stream
|
||||
SendAndClose(*RecordPromptUsageResponse) error
|
||||
}
|
||||
|
||||
type drpcRecorder_RecordPromptUsageStream struct {
|
||||
drpc.Stream
|
||||
}
|
||||
|
||||
func (x *drpcRecorder_RecordPromptUsageStream) SendAndClose(m *RecordPromptUsageResponse) error {
|
||||
if err := x.MsgSend(m, drpcEncoding_File_enterprise_x_aibridged_proto_aibridged_proto{}); err != nil {
|
||||
return err
|
||||
}
|
||||
return x.CloseSend()
|
||||
}
|
||||
|
||||
type DRPCRecorder_RecordToolUsageStream interface {
|
||||
drpc.Stream
|
||||
SendAndClose(*RecordToolUsageResponse) error
|
||||
}
|
||||
|
||||
type drpcRecorder_RecordToolUsageStream struct {
|
||||
drpc.Stream
|
||||
}
|
||||
|
||||
func (x *drpcRecorder_RecordToolUsageStream) SendAndClose(m *RecordToolUsageResponse) error {
|
||||
if err := x.MsgSend(m, drpcEncoding_File_enterprise_x_aibridged_proto_aibridged_proto{}); err != nil {
|
||||
return err
|
||||
}
|
||||
return x.CloseSend()
|
||||
}
|
||||
|
||||
type DRPCMCPConfiguratorClient interface {
|
||||
DRPCConn() drpc.Conn
|
||||
|
||||
GetMCPServerConfigs(ctx context.Context, in *GetMCPServerConfigsRequest) (*GetMCPServerConfigsResponse, error)
|
||||
GetMCPServerAccessTokensBatch(ctx context.Context, in *GetMCPServerAccessTokensBatchRequest) (*GetMCPServerAccessTokensBatchResponse, error)
|
||||
}
|
||||
|
||||
type drpcMCPConfiguratorClient struct {
|
||||
cc drpc.Conn
|
||||
}
|
||||
|
||||
func NewDRPCMCPConfiguratorClient(cc drpc.Conn) DRPCMCPConfiguratorClient {
|
||||
return &drpcMCPConfiguratorClient{cc}
|
||||
}
|
||||
|
||||
func (c *drpcMCPConfiguratorClient) DRPCConn() drpc.Conn { return c.cc }
|
||||
|
||||
func (c *drpcMCPConfiguratorClient) GetMCPServerConfigs(ctx context.Context, in *GetMCPServerConfigsRequest) (*GetMCPServerConfigsResponse, error) {
|
||||
out := new(GetMCPServerConfigsResponse)
|
||||
err := c.cc.Invoke(ctx, "/proto.MCPConfigurator/GetMCPServerConfigs", drpcEncoding_File_enterprise_x_aibridged_proto_aibridged_proto{}, in, out)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (c *drpcMCPConfiguratorClient) GetMCPServerAccessTokensBatch(ctx context.Context, in *GetMCPServerAccessTokensBatchRequest) (*GetMCPServerAccessTokensBatchResponse, error) {
|
||||
out := new(GetMCPServerAccessTokensBatchResponse)
|
||||
err := c.cc.Invoke(ctx, "/proto.MCPConfigurator/GetMCPServerAccessTokensBatch", drpcEncoding_File_enterprise_x_aibridged_proto_aibridged_proto{}, in, out)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
type DRPCMCPConfiguratorServer interface {
|
||||
GetMCPServerConfigs(context.Context, *GetMCPServerConfigsRequest) (*GetMCPServerConfigsResponse, error)
|
||||
GetMCPServerAccessTokensBatch(context.Context, *GetMCPServerAccessTokensBatchRequest) (*GetMCPServerAccessTokensBatchResponse, error)
|
||||
}
|
||||
|
||||
type DRPCMCPConfiguratorUnimplementedServer struct{}
|
||||
|
||||
func (s *DRPCMCPConfiguratorUnimplementedServer) GetMCPServerConfigs(context.Context, *GetMCPServerConfigsRequest) (*GetMCPServerConfigsResponse, error) {
|
||||
return nil, drpcerr.WithCode(errors.New("Unimplemented"), drpcerr.Unimplemented)
|
||||
}
|
||||
|
||||
func (s *DRPCMCPConfiguratorUnimplementedServer) GetMCPServerAccessTokensBatch(context.Context, *GetMCPServerAccessTokensBatchRequest) (*GetMCPServerAccessTokensBatchResponse, error) {
|
||||
return nil, drpcerr.WithCode(errors.New("Unimplemented"), drpcerr.Unimplemented)
|
||||
}
|
||||
|
||||
type DRPCMCPConfiguratorDescription struct{}
|
||||
|
||||
func (DRPCMCPConfiguratorDescription) NumMethods() int { return 2 }
|
||||
|
||||
func (DRPCMCPConfiguratorDescription) Method(n int) (string, drpc.Encoding, drpc.Receiver, interface{}, bool) {
|
||||
switch n {
|
||||
case 0:
|
||||
return "/proto.MCPConfigurator/GetMCPServerConfigs", drpcEncoding_File_enterprise_x_aibridged_proto_aibridged_proto{},
|
||||
func(srv interface{}, ctx context.Context, in1, in2 interface{}) (drpc.Message, error) {
|
||||
return srv.(DRPCMCPConfiguratorServer).
|
||||
GetMCPServerConfigs(
|
||||
ctx,
|
||||
in1.(*GetMCPServerConfigsRequest),
|
||||
)
|
||||
}, DRPCMCPConfiguratorServer.GetMCPServerConfigs, true
|
||||
case 1:
|
||||
return "/proto.MCPConfigurator/GetMCPServerAccessTokensBatch", drpcEncoding_File_enterprise_x_aibridged_proto_aibridged_proto{},
|
||||
func(srv interface{}, ctx context.Context, in1, in2 interface{}) (drpc.Message, error) {
|
||||
return srv.(DRPCMCPConfiguratorServer).
|
||||
GetMCPServerAccessTokensBatch(
|
||||
ctx,
|
||||
in1.(*GetMCPServerAccessTokensBatchRequest),
|
||||
)
|
||||
}, DRPCMCPConfiguratorServer.GetMCPServerAccessTokensBatch, true
|
||||
default:
|
||||
return "", nil, nil, nil, false
|
||||
}
|
||||
}
|
||||
|
||||
func DRPCRegisterMCPConfigurator(mux drpc.Mux, impl DRPCMCPConfiguratorServer) error {
|
||||
return mux.Register(impl, DRPCMCPConfiguratorDescription{})
|
||||
}
|
||||
|
||||
type DRPCMCPConfigurator_GetMCPServerConfigsStream interface {
|
||||
drpc.Stream
|
||||
SendAndClose(*GetMCPServerConfigsResponse) error
|
||||
}
|
||||
|
||||
type drpcMCPConfigurator_GetMCPServerConfigsStream struct {
|
||||
drpc.Stream
|
||||
}
|
||||
|
||||
func (x *drpcMCPConfigurator_GetMCPServerConfigsStream) SendAndClose(m *GetMCPServerConfigsResponse) error {
|
||||
if err := x.MsgSend(m, drpcEncoding_File_enterprise_x_aibridged_proto_aibridged_proto{}); err != nil {
|
||||
return err
|
||||
}
|
||||
return x.CloseSend()
|
||||
}
|
||||
|
||||
type DRPCMCPConfigurator_GetMCPServerAccessTokensBatchStream interface {
|
||||
drpc.Stream
|
||||
SendAndClose(*GetMCPServerAccessTokensBatchResponse) error
|
||||
}
|
||||
|
||||
type drpcMCPConfigurator_GetMCPServerAccessTokensBatchStream struct {
|
||||
drpc.Stream
|
||||
}
|
||||
|
||||
func (x *drpcMCPConfigurator_GetMCPServerAccessTokensBatchStream) SendAndClose(m *GetMCPServerAccessTokensBatchResponse) error {
|
||||
if err := x.MsgSend(m, drpcEncoding_File_enterprise_x_aibridged_proto_aibridged_proto{}); err != nil {
|
||||
return err
|
||||
}
|
||||
return x.CloseSend()
|
||||
}
|
||||
|
||||
type DRPCAuthorizerClient interface {
|
||||
DRPCConn() drpc.Conn
|
||||
|
||||
IsAuthorized(ctx context.Context, in *IsAuthorizedRequest) (*IsAuthorizedResponse, error)
|
||||
}
|
||||
|
||||
type drpcAuthorizerClient struct {
|
||||
cc drpc.Conn
|
||||
}
|
||||
|
||||
func NewDRPCAuthorizerClient(cc drpc.Conn) DRPCAuthorizerClient {
|
||||
return &drpcAuthorizerClient{cc}
|
||||
}
|
||||
|
||||
func (c *drpcAuthorizerClient) DRPCConn() drpc.Conn { return c.cc }
|
||||
|
||||
func (c *drpcAuthorizerClient) IsAuthorized(ctx context.Context, in *IsAuthorizedRequest) (*IsAuthorizedResponse, error) {
|
||||
out := new(IsAuthorizedResponse)
|
||||
err := c.cc.Invoke(ctx, "/proto.Authorizer/IsAuthorized", drpcEncoding_File_enterprise_x_aibridged_proto_aibridged_proto{}, in, out)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
type DRPCAuthorizerServer interface {
|
||||
IsAuthorized(context.Context, *IsAuthorizedRequest) (*IsAuthorizedResponse, error)
|
||||
}
|
||||
|
||||
type DRPCAuthorizerUnimplementedServer struct{}
|
||||
|
||||
func (s *DRPCAuthorizerUnimplementedServer) IsAuthorized(context.Context, *IsAuthorizedRequest) (*IsAuthorizedResponse, error) {
|
||||
return nil, drpcerr.WithCode(errors.New("Unimplemented"), drpcerr.Unimplemented)
|
||||
}
|
||||
|
||||
type DRPCAuthorizerDescription struct{}
|
||||
|
||||
func (DRPCAuthorizerDescription) NumMethods() int { return 1 }
|
||||
|
||||
func (DRPCAuthorizerDescription) Method(n int) (string, drpc.Encoding, drpc.Receiver, interface{}, bool) {
|
||||
switch n {
|
||||
case 0:
|
||||
return "/proto.Authorizer/IsAuthorized", drpcEncoding_File_enterprise_x_aibridged_proto_aibridged_proto{},
|
||||
func(srv interface{}, ctx context.Context, in1, in2 interface{}) (drpc.Message, error) {
|
||||
return srv.(DRPCAuthorizerServer).
|
||||
IsAuthorized(
|
||||
ctx,
|
||||
in1.(*IsAuthorizedRequest),
|
||||
)
|
||||
}, DRPCAuthorizerServer.IsAuthorized, true
|
||||
default:
|
||||
return "", nil, nil, nil, false
|
||||
}
|
||||
}
|
||||
|
||||
func DRPCRegisterAuthorizer(mux drpc.Mux, impl DRPCAuthorizerServer) error {
|
||||
return mux.Register(impl, DRPCAuthorizerDescription{})
|
||||
}
|
||||
|
||||
type DRPCAuthorizer_IsAuthorizedStream interface {
|
||||
drpc.Stream
|
||||
SendAndClose(*IsAuthorizedResponse) error
|
||||
}
|
||||
|
||||
type drpcAuthorizer_IsAuthorizedStream struct {
|
||||
drpc.Stream
|
||||
}
|
||||
|
||||
func (x *drpcAuthorizer_IsAuthorizedStream) SendAndClose(m *IsAuthorizedResponse) error {
|
||||
if err := x.MsgSend(m, drpcEncoding_File_enterprise_x_aibridged_proto_aibridged_proto{}); err != nil {
|
||||
return err
|
||||
}
|
||||
return x.CloseSend()
|
||||
}
|
||||
@@ -0,0 +1,430 @@
|
||||
package aibridgedserver
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"crypto/subtle"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"net/url"
|
||||
"slices"
|
||||
"sync"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/hashicorp/go-multierror"
|
||||
"golang.org/x/xerrors"
|
||||
"google.golang.org/protobuf/types/known/anypb"
|
||||
"google.golang.org/protobuf/types/known/structpb"
|
||||
|
||||
"cdr.dev/slog"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/dbauthz"
|
||||
"github.com/coder/coder/v2/coderd/database/dbtime"
|
||||
"github.com/coder/coder/v2/coderd/externalauth"
|
||||
"github.com/coder/coder/v2/coderd/httpmw"
|
||||
codermcp "github.com/coder/coder/v2/coderd/mcp"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/enterprise/x/aibridged/proto"
|
||||
)
|
||||
|
||||
var (
|
||||
ErrExpiredOrInvalidOAuthToken = xerrors.New("expired or invalid OAuth2 token")
|
||||
ErrNoMCPConfigFound = xerrors.New("no MCP config found")
|
||||
|
||||
// These errors are returned by IsAuthorized. Since they're just returned as
|
||||
// a generic dRPC error, it's difficult to tell them apart without string
|
||||
// matching.
|
||||
// TODO: return these errors to the client in a more structured/comparable
|
||||
// way.
|
||||
ErrInvalidKey = xerrors.New("invalid key")
|
||||
ErrUnknownKey = xerrors.New("unknown key")
|
||||
ErrExpired = xerrors.New("expired")
|
||||
ErrUnknownUser = xerrors.New("unknown user")
|
||||
ErrDeletedUser = xerrors.New("deleted user")
|
||||
ErrSystemUser = xerrors.New("system user")
|
||||
|
||||
ErrNoExternalAuthLinkFound = xerrors.New("no external auth link found")
|
||||
)
|
||||
|
||||
var (
|
||||
_ proto.DRPCAuthorizerServer = &Server{}
|
||||
_ proto.DRPCMCPConfiguratorServer = &Server{}
|
||||
_ proto.DRPCRecorderServer = &Server{}
|
||||
)
|
||||
|
||||
type store interface {
|
||||
// Recorder-related queries.
|
||||
InsertAIBridgeInterception(ctx context.Context, arg database.InsertAIBridgeInterceptionParams) (database.AIBridgeInterception, error)
|
||||
InsertAIBridgeTokenUsage(ctx context.Context, arg database.InsertAIBridgeTokenUsageParams) error
|
||||
InsertAIBridgeUserPrompt(ctx context.Context, arg database.InsertAIBridgeUserPromptParams) error
|
||||
InsertAIBridgeToolUsage(ctx context.Context, arg database.InsertAIBridgeToolUsageParams) error
|
||||
|
||||
// MCPConfigurator-related queries.
|
||||
GetExternalAuthLinksByUserID(ctx context.Context, userID uuid.UUID) ([]database.ExternalAuthLink, error)
|
||||
|
||||
// Authorizer-related queries.
|
||||
GetAPIKeyByID(ctx context.Context, id string) (database.APIKey, error)
|
||||
GetUserByID(ctx context.Context, id uuid.UUID) (database.User, error)
|
||||
}
|
||||
|
||||
type Server struct {
|
||||
// lifecycleCtx must be tied to the API server's lifecycle
|
||||
// as when the API server shuts down, we want to cancel any
|
||||
// long-running operations.
|
||||
lifecycleCtx context.Context
|
||||
store store
|
||||
logger slog.Logger
|
||||
externalAuthConfigs map[string]*externalauth.Config
|
||||
|
||||
coderMCPConfig *proto.MCPServerConfig // may be nil if not available
|
||||
}
|
||||
|
||||
func NewServer(lifecycleCtx context.Context, store store, logger slog.Logger, accessURL string, externalAuthConfigs []*externalauth.Config, experiments codersdk.Experiments) (*Server, error) {
|
||||
eac := make(map[string]*externalauth.Config, len(externalAuthConfigs))
|
||||
|
||||
for _, cfg := range externalAuthConfigs {
|
||||
// Only External Auth configs which are configured with an MCP URL are relevant to aibridged.
|
||||
if cfg.MCPURL == "" {
|
||||
continue
|
||||
}
|
||||
eac[cfg.ID] = cfg
|
||||
}
|
||||
|
||||
coderMCPConfig, err := getCoderMCPServerConfig(experiments, accessURL)
|
||||
if err != nil {
|
||||
logger.Warn(lifecycleCtx, "failed to retrieve coder MCP server config, Coder MCP will not be available", slog.Error(err))
|
||||
}
|
||||
|
||||
return &Server{
|
||||
lifecycleCtx: lifecycleCtx,
|
||||
store: store,
|
||||
logger: logger.Named("aibridgedserver"),
|
||||
externalAuthConfigs: eac,
|
||||
coderMCPConfig: coderMCPConfig,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *Server) RecordInterception(ctx context.Context, in *proto.RecordInterceptionRequest) (*proto.RecordInterceptionResponse, error) {
|
||||
//nolint:gocritic // AIBridged has specific authz rules.
|
||||
ctx = dbauthz.AsAIBridged(ctx)
|
||||
|
||||
intcID, err := uuid.Parse(in.GetId())
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("invalid interception ID %q: %w", in.GetId(), err)
|
||||
}
|
||||
initID, err := uuid.Parse(in.GetInitiatorId())
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("invalid initiator ID %q: %w", in.GetInitiatorId(), err)
|
||||
}
|
||||
|
||||
_, err = s.store.InsertAIBridgeInterception(ctx, database.InsertAIBridgeInterceptionParams{
|
||||
ID: intcID,
|
||||
InitiatorID: initID,
|
||||
Provider: in.Provider,
|
||||
Model: in.Model,
|
||||
Metadata: marshalMetadata(ctx, s.logger, in.GetMetadata()),
|
||||
StartedAt: in.StartedAt.AsTime(),
|
||||
})
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("start interception: %w", err)
|
||||
}
|
||||
|
||||
return &proto.RecordInterceptionResponse{}, nil
|
||||
}
|
||||
|
||||
func (s *Server) RecordTokenUsage(ctx context.Context, in *proto.RecordTokenUsageRequest) (*proto.RecordTokenUsageResponse, error) {
|
||||
//nolint:gocritic // AIBridged has specific authz rules.
|
||||
ctx = dbauthz.AsAIBridged(ctx)
|
||||
|
||||
intcID, err := uuid.Parse(in.GetInterceptionId())
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("failed to parse interception_id %q: %w", in.GetInterceptionId(), err)
|
||||
}
|
||||
|
||||
err = s.store.InsertAIBridgeTokenUsage(ctx, database.InsertAIBridgeTokenUsageParams{
|
||||
ID: uuid.New(),
|
||||
InterceptionID: intcID,
|
||||
ProviderResponseID: in.GetMsgId(),
|
||||
InputTokens: in.GetInputTokens(),
|
||||
OutputTokens: in.GetOutputTokens(),
|
||||
Metadata: marshalMetadata(ctx, s.logger, in.GetMetadata()),
|
||||
CreatedAt: in.GetCreatedAt().AsTime(),
|
||||
})
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("insert token usage: %w", err)
|
||||
}
|
||||
return &proto.RecordTokenUsageResponse{}, nil
|
||||
}
|
||||
|
||||
func (s *Server) RecordPromptUsage(ctx context.Context, in *proto.RecordPromptUsageRequest) (*proto.RecordPromptUsageResponse, error) {
|
||||
//nolint:gocritic // AIBridged has specific authz rules.
|
||||
ctx = dbauthz.AsAIBridged(ctx)
|
||||
|
||||
intcID, err := uuid.Parse(in.GetInterceptionId())
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("failed to parse interception_id %q: %w", in.GetInterceptionId(), err)
|
||||
}
|
||||
|
||||
err = s.store.InsertAIBridgeUserPrompt(ctx, database.InsertAIBridgeUserPromptParams{
|
||||
ID: uuid.New(),
|
||||
InterceptionID: intcID,
|
||||
ProviderResponseID: in.GetMsgId(),
|
||||
Prompt: in.GetPrompt(),
|
||||
Metadata: marshalMetadata(ctx, s.logger, in.GetMetadata()),
|
||||
CreatedAt: in.GetCreatedAt().AsTime(),
|
||||
})
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("insert user prompt: %w", err)
|
||||
}
|
||||
return &proto.RecordPromptUsageResponse{}, nil
|
||||
}
|
||||
|
||||
func (s *Server) RecordToolUsage(ctx context.Context, in *proto.RecordToolUsageRequest) (*proto.RecordToolUsageResponse, error) {
|
||||
//nolint:gocritic // AIBridged has specific authz rules.
|
||||
ctx = dbauthz.AsAIBridged(ctx)
|
||||
|
||||
intcID, err := uuid.Parse(in.GetInterceptionId())
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("failed to parse interception_id %q: %w", in.GetInterceptionId(), err)
|
||||
}
|
||||
|
||||
err = s.store.InsertAIBridgeToolUsage(ctx, database.InsertAIBridgeToolUsageParams{
|
||||
ID: uuid.New(),
|
||||
InterceptionID: intcID,
|
||||
ProviderResponseID: in.GetMsgId(),
|
||||
ServerUrl: sql.NullString{String: in.GetServerUrl(), Valid: in.ServerUrl != nil},
|
||||
Tool: in.GetTool(),
|
||||
Input: in.GetInput(),
|
||||
Injected: in.GetInjected(),
|
||||
InvocationError: sql.NullString{String: in.GetInvocationError(), Valid: in.InvocationError != nil},
|
||||
Metadata: marshalMetadata(ctx, s.logger, in.GetMetadata()),
|
||||
CreatedAt: in.GetCreatedAt().AsTime(),
|
||||
})
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("insert tool usage: %w", err)
|
||||
}
|
||||
return &proto.RecordToolUsageResponse{}, nil
|
||||
}
|
||||
|
||||
func (s *Server) GetMCPServerConfigs(_ context.Context, _ *proto.GetMCPServerConfigsRequest) (*proto.GetMCPServerConfigsResponse, error) {
|
||||
cfgs := make([]*proto.MCPServerConfig, 0, len(s.externalAuthConfigs))
|
||||
for _, eac := range s.externalAuthConfigs {
|
||||
var allowlist, denylist string
|
||||
if eac.MCPToolAllowRegex != nil {
|
||||
allowlist = eac.MCPToolAllowRegex.String()
|
||||
}
|
||||
if eac.MCPToolDenyRegex != nil {
|
||||
denylist = eac.MCPToolDenyRegex.String()
|
||||
}
|
||||
|
||||
cfgs = append(cfgs, &proto.MCPServerConfig{
|
||||
Id: eac.ID,
|
||||
Url: eac.MCPURL,
|
||||
ToolAllowRegex: allowlist,
|
||||
ToolDenyRegex: denylist,
|
||||
})
|
||||
}
|
||||
|
||||
return &proto.GetMCPServerConfigsResponse{
|
||||
CoderMcpConfig: s.coderMCPConfig, // it's fine if this is nil
|
||||
ExternalAuthMcpConfigs: cfgs,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *Server) GetMCPServerAccessTokensBatch(ctx context.Context, in *proto.GetMCPServerAccessTokensBatchRequest) (*proto.GetMCPServerAccessTokensBatchResponse, error) {
|
||||
if len(in.GetMcpServerConfigIds()) == 0 {
|
||||
return &proto.GetMCPServerAccessTokensBatchResponse{}, nil
|
||||
}
|
||||
|
||||
userID, err := uuid.Parse(in.GetUserId())
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("parse user_id: %w", err)
|
||||
}
|
||||
|
||||
//nolint:gocritic // AIBridged has specific authz rules.
|
||||
ctx = dbauthz.AsAIBridged(ctx)
|
||||
links, err := s.store.GetExternalAuthLinksByUserID(ctx, userID)
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("fetch external auth links: %w", err)
|
||||
}
|
||||
|
||||
if len(links) == 0 {
|
||||
return &proto.GetMCPServerAccessTokensBatchResponse{}, nil
|
||||
}
|
||||
|
||||
// Ensure unique to prevent unnecessary effort.
|
||||
ids := in.GetMcpServerConfigIds()
|
||||
slices.Sort(ids)
|
||||
ids = slices.Compact(ids)
|
||||
|
||||
var (
|
||||
wg sync.WaitGroup
|
||||
errs error
|
||||
|
||||
mu sync.Mutex
|
||||
tokens = make(map[string]string, len(ids))
|
||||
tokenErrs = make(map[string]string)
|
||||
)
|
||||
|
||||
externalAuthLoop:
|
||||
for _, id := range ids {
|
||||
eac, ok := s.externalAuthConfigs[id]
|
||||
if !ok {
|
||||
mu.Lock()
|
||||
s.logger.Warn(ctx, "no MCP server config found by given ID", slog.F("id", id))
|
||||
tokenErrs[id] = ErrNoMCPConfigFound.Error()
|
||||
mu.Unlock()
|
||||
continue
|
||||
}
|
||||
|
||||
for _, link := range links {
|
||||
if link.ProviderID != eac.ID {
|
||||
continue
|
||||
}
|
||||
|
||||
// Validate all configured External Auth links concurrently.
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
|
||||
// TODO: timeout.
|
||||
valid, _, validateErr := eac.ValidateToken(ctx, link.OAuthToken())
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
if !valid {
|
||||
// TODO: attempt refresh.
|
||||
s.logger.Warn(ctx, "invalid/expired access token, cannot auto-configure MCP", slog.F("provider", link.ProviderID), slog.Error(validateErr))
|
||||
tokenErrs[id] = ErrExpiredOrInvalidOAuthToken.Error()
|
||||
return
|
||||
}
|
||||
|
||||
if validateErr != nil {
|
||||
errs = multierror.Append(errs, validateErr)
|
||||
tokenErrs[id] = validateErr.Error()
|
||||
} else {
|
||||
tokens[id] = link.OAuthAccessToken
|
||||
}
|
||||
}()
|
||||
|
||||
continue externalAuthLoop
|
||||
}
|
||||
|
||||
// No link found for this external auth config, so include a generic
|
||||
// error.
|
||||
mu.Lock()
|
||||
tokenErrs[id] = ErrNoExternalAuthLinkFound.Error()
|
||||
mu.Unlock()
|
||||
}
|
||||
|
||||
wg.Wait()
|
||||
return &proto.GetMCPServerAccessTokensBatchResponse{
|
||||
AccessTokens: tokens,
|
||||
Errors: tokenErrs,
|
||||
}, errs
|
||||
}
|
||||
|
||||
// IsAuthorized validates a given Coder API key and returns the user ID to which it belongs (if valid).
|
||||
//
|
||||
// NOTE: this should really be using the code from [httpmw.ExtractAPIKey]. That function not only validates the key
|
||||
// but handles many other cases like updating last used, expiry, etc. This code does not currently use it for
|
||||
// a few reasons:
|
||||
//
|
||||
// 1. [httpmw.ExtractAPIKey] relies on keys being given in specific headers [httpmw.APITokenFromRequest] which AI
|
||||
// bridge requests will not conform to.
|
||||
// 2. The code mixes many different concerns, and handles HTTP responses too, which is undesirable here.
|
||||
// 3. The core logic would need to be extracted, but that will surely be a complex & time-consuming distraction right now.
|
||||
// 4. Once we have an Early Access release of AI Bridge, we need to return to this.
|
||||
//
|
||||
// TODO: replace with logic from [httpmw.ExtractAPIKey].
|
||||
func (s *Server) IsAuthorized(ctx context.Context, in *proto.IsAuthorizedRequest) (*proto.IsAuthorizedResponse, error) {
|
||||
//nolint:gocritic // AIBridged has specific authz rules.
|
||||
ctx = dbauthz.AsAIBridged(ctx)
|
||||
|
||||
// Key matches expected format.
|
||||
keyID, keySecret, err := httpmw.SplitAPIToken(in.GetKey())
|
||||
if err != nil {
|
||||
return nil, ErrInvalidKey
|
||||
}
|
||||
|
||||
// Key exists.
|
||||
key, err := s.store.GetAPIKeyByID(ctx, keyID)
|
||||
if err != nil {
|
||||
s.logger.Warn(ctx, "failed to retrieve API key by id", slog.F("key_id", keyID), slog.Error(err))
|
||||
return nil, ErrUnknownKey
|
||||
}
|
||||
|
||||
// Key has not expired.
|
||||
now := dbtime.Now()
|
||||
if key.ExpiresAt.Before(now) {
|
||||
return nil, ErrExpired
|
||||
}
|
||||
|
||||
// Key secret matches.
|
||||
hashedSecret := sha256.Sum256([]byte(keySecret))
|
||||
if subtle.ConstantTimeCompare(key.HashedSecret, hashedSecret[:]) != 1 {
|
||||
return nil, ErrInvalidKey
|
||||
}
|
||||
|
||||
// User exists.
|
||||
user, err := s.store.GetUserByID(ctx, key.UserID)
|
||||
if err != nil {
|
||||
s.logger.Warn(ctx, "failed to retrieve API key user", slog.F("key_id", keyID), slog.F("user_id", key.UserID), slog.Error(err))
|
||||
return nil, ErrUnknownUser
|
||||
}
|
||||
|
||||
// User is not deleted or a system user.
|
||||
if user.Deleted {
|
||||
return nil, ErrDeletedUser
|
||||
}
|
||||
if user.IsSystem {
|
||||
return nil, ErrSystemUser
|
||||
}
|
||||
|
||||
return &proto.IsAuthorizedResponse{
|
||||
OwnerId: key.UserID.String(),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func getCoderMCPServerConfig(experiments codersdk.Experiments, accessURL string) (*proto.MCPServerConfig, error) {
|
||||
// Both the MCP & OAuth2 experiments are currently required in order to use our
|
||||
// internal MCP server.
|
||||
if !experiments.Enabled(codersdk.ExperimentMCPServerHTTP) {
|
||||
return nil, xerrors.Errorf("%q experiment not enabled", codersdk.ExperimentMCPServerHTTP)
|
||||
}
|
||||
if !experiments.Enabled(codersdk.ExperimentOAuth2) {
|
||||
return nil, xerrors.Errorf("%q experiment not enabled", codersdk.ExperimentOAuth2)
|
||||
}
|
||||
|
||||
u, err := url.JoinPath(accessURL, codermcp.MCPEndpoint)
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("build MCP URL with %q: %w", accessURL, err)
|
||||
}
|
||||
|
||||
return &proto.MCPServerConfig{
|
||||
Id: "coder",
|
||||
Url: u,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// marshalMetadata attempts to marshal the given metadata map into a
|
||||
// JSON-encoded byte slice. If the marshaling fails, the function logs a
|
||||
// warning and returns nil. The supplied context is only used for logging.
|
||||
func marshalMetadata(ctx context.Context, logger slog.Logger, in map[string]*anypb.Any) []byte {
|
||||
mdMap := make(map[string]any, len(in))
|
||||
for k, v := range in {
|
||||
if v == nil {
|
||||
continue
|
||||
}
|
||||
var sv structpb.Value
|
||||
if err := v.UnmarshalTo(&sv); err == nil {
|
||||
mdMap[k] = sv.AsInterface()
|
||||
}
|
||||
}
|
||||
out, err := json.Marshal(mdMap)
|
||||
if err != nil {
|
||||
logger.Warn(ctx, "failed to marshal aibridge metadata from proto to JSON", slog.F("metadata", in), slog.Error(err))
|
||||
return nil
|
||||
}
|
||||
return out
|
||||
}
|
||||
@@ -0,0 +1,88 @@
|
||||
package aibridgedserver
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"math"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/protobuf/proto"
|
||||
"google.golang.org/protobuf/types/known/anypb"
|
||||
"google.golang.org/protobuf/types/known/structpb"
|
||||
|
||||
"cdr.dev/slog"
|
||||
"cdr.dev/slog/sloggers/slogtest"
|
||||
)
|
||||
|
||||
func TestMarshalMetadata(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("NilData", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}).Leveled(slog.LevelDebug)
|
||||
out := marshalMetadata(context.Background(), logger, nil)
|
||||
require.JSONEq(t, "{}", string(out))
|
||||
})
|
||||
|
||||
t.Run("WithData", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}).Leveled(slog.LevelDebug)
|
||||
|
||||
list := structpb.NewListValue(&structpb.ListValue{Values: []*structpb.Value{
|
||||
structpb.NewStringValue("a"),
|
||||
structpb.NewNumberValue(1),
|
||||
structpb.NewBoolValue(false),
|
||||
}})
|
||||
obj := structpb.NewStructValue(&structpb.Struct{Fields: map[string]*structpb.Value{
|
||||
"a": structpb.NewStringValue("b"),
|
||||
"n": structpb.NewNumberValue(3),
|
||||
}})
|
||||
|
||||
nonValue := mustMarshalAny(t, &structpb.Struct{Fields: map[string]*structpb.Value{
|
||||
"ignored": structpb.NewStringValue("yes"),
|
||||
}})
|
||||
invalid := &anypb.Any{TypeUrl: "type.googleapis.com/google.protobuf.Value", Value: []byte{0xff, 0x00}}
|
||||
|
||||
in := map[string]*anypb.Any{
|
||||
"null": mustMarshalAny(t, structpb.NewNullValue()),
|
||||
// Scalars
|
||||
"string": mustMarshalAny(t, structpb.NewStringValue("hello")),
|
||||
"bool": mustMarshalAny(t, structpb.NewBoolValue(true)),
|
||||
"number": mustMarshalAny(t, structpb.NewNumberValue(42)),
|
||||
// Complex types
|
||||
"list": mustMarshalAny(t, list),
|
||||
"object": mustMarshalAny(t, obj),
|
||||
// Extra valid entries
|
||||
"ok": mustMarshalAny(t, structpb.NewStringValue("present")),
|
||||
"nan": mustMarshalAny(t, structpb.NewNumberValue(math.NaN())),
|
||||
// Entries that should be ignored
|
||||
"invalid": invalid,
|
||||
"non_value": nonValue,
|
||||
}
|
||||
|
||||
out := marshalMetadata(context.Background(), logger, in)
|
||||
require.NotNil(t, out)
|
||||
var got map[string]any
|
||||
require.NoError(t, json.Unmarshal(out, &got))
|
||||
|
||||
expected := map[string]any{
|
||||
"string": "hello",
|
||||
"bool": true,
|
||||
"number": float64(42),
|
||||
"null": nil,
|
||||
"list": []any{"a", float64(1), false},
|
||||
"object": map[string]any{"a": "b", "n": float64(3)},
|
||||
"ok": "present",
|
||||
"nan": "NaN",
|
||||
}
|
||||
require.Equal(t, expected, got)
|
||||
})
|
||||
}
|
||||
|
||||
func mustMarshalAny(t testing.TB, m proto.Message) *anypb.Any {
|
||||
t.Helper()
|
||||
a, err := anypb.New(m)
|
||||
require.NoError(t, err)
|
||||
return a
|
||||
}
|
||||
@@ -0,0 +1,710 @@
|
||||
package aibridgedserver_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/url"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/sqlc-dev/pqtype"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/mock/gomock"
|
||||
protobufproto "google.golang.org/protobuf/proto"
|
||||
"google.golang.org/protobuf/types/known/anypb"
|
||||
"google.golang.org/protobuf/types/known/structpb"
|
||||
"google.golang.org/protobuf/types/known/timestamppb"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/dbmock"
|
||||
"github.com/coder/coder/v2/coderd/database/dbtime"
|
||||
"github.com/coder/coder/v2/coderd/externalauth"
|
||||
codermcp "github.com/coder/coder/v2/coderd/mcp"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/cryptorand"
|
||||
"github.com/coder/coder/v2/enterprise/x/aibridged/proto"
|
||||
"github.com/coder/coder/v2/enterprise/x/aibridgedserver"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
)
|
||||
|
||||
var requiredExperiments = []codersdk.Experiment{
|
||||
codersdk.ExperimentMCPServerHTTP, codersdk.ExperimentOAuth2,
|
||||
}
|
||||
|
||||
// TestAuthorization validates the authorization logic.
|
||||
// No other tests are explicitly defined in this package because aibridgedserver is
|
||||
// tested via integration tests in the aibridged package (see aibridged/aibridged_integration_test.go).
|
||||
func TestAuthorization(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
// Key will be set to the same key passed to mocksFn if unset.
|
||||
key string
|
||||
// mocksFn is called with a valid API key and user. If the test needs
|
||||
// invalid values, it should just mutate them directly.
|
||||
mocksFn func(db *dbmock.MockStore, apiKey database.APIKey, user database.User)
|
||||
expectedErr error
|
||||
}{
|
||||
{
|
||||
name: "invalid key format",
|
||||
key: "foo",
|
||||
expectedErr: aibridgedserver.ErrInvalidKey,
|
||||
},
|
||||
{
|
||||
name: "unknown key",
|
||||
expectedErr: aibridgedserver.ErrUnknownKey,
|
||||
mocksFn: func(db *dbmock.MockStore, apiKey database.APIKey, user database.User) {
|
||||
db.EXPECT().GetAPIKeyByID(gomock.Any(), apiKey.ID).Times(1).Return(database.APIKey{}, sql.ErrNoRows)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "expired",
|
||||
expectedErr: aibridgedserver.ErrExpired,
|
||||
mocksFn: func(db *dbmock.MockStore, apiKey database.APIKey, user database.User) {
|
||||
apiKey.ExpiresAt = dbtime.Now().Add(-time.Hour)
|
||||
db.EXPECT().GetAPIKeyByID(gomock.Any(), apiKey.ID).Times(1).Return(apiKey, nil)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "invalid key secret",
|
||||
expectedErr: aibridgedserver.ErrInvalidKey,
|
||||
mocksFn: func(db *dbmock.MockStore, apiKey database.APIKey, user database.User) {
|
||||
apiKey.HashedSecret = []byte("differentsecret")
|
||||
db.EXPECT().GetAPIKeyByID(gomock.Any(), apiKey.ID).Times(1).Return(apiKey, nil)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "unknown user",
|
||||
expectedErr: aibridgedserver.ErrUnknownUser,
|
||||
mocksFn: func(db *dbmock.MockStore, apiKey database.APIKey, user database.User) {
|
||||
db.EXPECT().GetAPIKeyByID(gomock.Any(), apiKey.ID).Times(1).Return(apiKey, nil)
|
||||
db.EXPECT().GetUserByID(gomock.Any(), user.ID).Times(1).Return(database.User{}, sql.ErrNoRows)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "deleted user",
|
||||
expectedErr: aibridgedserver.ErrDeletedUser,
|
||||
mocksFn: func(db *dbmock.MockStore, apiKey database.APIKey, user database.User) {
|
||||
db.EXPECT().GetAPIKeyByID(gomock.Any(), apiKey.ID).Times(1).Return(apiKey, nil)
|
||||
db.EXPECT().GetUserByID(gomock.Any(), user.ID).Times(1).Return(database.User{ID: user.ID, Deleted: true}, nil)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "system user",
|
||||
expectedErr: aibridgedserver.ErrSystemUser,
|
||||
mocksFn: func(db *dbmock.MockStore, apiKey database.APIKey, user database.User) {
|
||||
db.EXPECT().GetAPIKeyByID(gomock.Any(), apiKey.ID).Times(1).Return(apiKey, nil)
|
||||
db.EXPECT().GetUserByID(gomock.Any(), user.ID).Times(1).Return(database.User{ID: user.ID, IsSystem: true}, nil)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "valid",
|
||||
mocksFn: func(db *dbmock.MockStore, apiKey database.APIKey, user database.User) {
|
||||
db.EXPECT().GetAPIKeyByID(gomock.Any(), apiKey.ID).Times(1).Return(apiKey, nil)
|
||||
db.EXPECT().GetUserByID(gomock.Any(), user.ID).Times(1).Return(user, nil)
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctrl := gomock.NewController(t)
|
||||
db := dbmock.NewMockStore(ctrl)
|
||||
logger := testutil.Logger(t)
|
||||
|
||||
// Make a fake user and an API key for the mock calls.
|
||||
now := dbtime.Now()
|
||||
user := database.User{
|
||||
ID: uuid.New(),
|
||||
Email: "test@coder.com",
|
||||
Username: "test",
|
||||
Name: "Test User",
|
||||
CreatedAt: now,
|
||||
UpdatedAt: now,
|
||||
RBACRoles: []string{},
|
||||
LoginType: database.LoginTypePassword,
|
||||
Status: database.UserStatusActive,
|
||||
LastSeenAt: now,
|
||||
}
|
||||
|
||||
keyID, _ := cryptorand.String(10)
|
||||
keySecret, _ := cryptorand.String(22)
|
||||
token := fmt.Sprintf("%s-%s", keyID, keySecret)
|
||||
keySecretHashed := sha256.Sum256([]byte(keySecret))
|
||||
apiKey := database.APIKey{
|
||||
ID: keyID,
|
||||
LifetimeSeconds: 86400, // default in db
|
||||
HashedSecret: keySecretHashed[:],
|
||||
IPAddress: pqtype.Inet{
|
||||
IPNet: net.IPNet{
|
||||
IP: net.IPv4(127, 0, 0, 1),
|
||||
Mask: net.IPv4Mask(255, 255, 255, 255),
|
||||
},
|
||||
Valid: true,
|
||||
},
|
||||
UserID: user.ID,
|
||||
LastUsed: now,
|
||||
ExpiresAt: now.Add(time.Hour),
|
||||
CreatedAt: now,
|
||||
UpdatedAt: now,
|
||||
LoginType: database.LoginTypePassword,
|
||||
Scopes: []database.APIKeyScope{database.APIKeyScopeAll},
|
||||
TokenName: "",
|
||||
}
|
||||
if tc.key == "" {
|
||||
tc.key = token
|
||||
}
|
||||
|
||||
// Define any case-specific mocks.
|
||||
if tc.mocksFn != nil {
|
||||
tc.mocksFn(db, apiKey, user)
|
||||
}
|
||||
|
||||
srv, err := aibridgedserver.NewServer(t.Context(), db, logger, "/", nil, requiredExperiments)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, srv)
|
||||
|
||||
_, err = srv.IsAuthorized(t.Context(), &proto.IsAuthorizedRequest{Key: tc.key})
|
||||
if tc.expectedErr != nil {
|
||||
require.Error(t, err)
|
||||
require.ErrorIs(t, err, tc.expectedErr)
|
||||
} else {
|
||||
require.NoError(t, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetMCPServerConfigs(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
externalAuthCfgs := []*externalauth.Config{
|
||||
{
|
||||
ID: "1",
|
||||
MCPURL: "1.com/mcp",
|
||||
},
|
||||
{
|
||||
ID: "2", // Will not be eligible for inclusion since MCPURL is not defined.
|
||||
},
|
||||
}
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
experiments codersdk.Experiments
|
||||
externalAuthConfigs []*externalauth.Config
|
||||
expectCoderMCP bool
|
||||
expectedExternalMCP bool
|
||||
}{
|
||||
{
|
||||
name: "experiments not enabled",
|
||||
experiments: codersdk.Experiments{},
|
||||
},
|
||||
{
|
||||
name: "MCP experiment enabled, not OAuth2",
|
||||
experiments: codersdk.Experiments{codersdk.ExperimentMCPServerHTTP},
|
||||
},
|
||||
{
|
||||
name: "OAuth2 experiment enabled, not MCP",
|
||||
experiments: codersdk.Experiments{codersdk.ExperimentOAuth2},
|
||||
},
|
||||
{
|
||||
name: "only internal MCP",
|
||||
experiments: requiredExperiments,
|
||||
expectCoderMCP: true,
|
||||
},
|
||||
{
|
||||
name: "only external MCP",
|
||||
externalAuthConfigs: externalAuthCfgs,
|
||||
expectedExternalMCP: true,
|
||||
},
|
||||
{
|
||||
name: "both internal & external MCP",
|
||||
experiments: requiredExperiments,
|
||||
externalAuthConfigs: externalAuthCfgs,
|
||||
expectCoderMCP: true,
|
||||
expectedExternalMCP: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctrl := gomock.NewController(t)
|
||||
db := dbmock.NewMockStore(ctrl)
|
||||
logger := testutil.Logger(t)
|
||||
|
||||
accessURL := "https://my-cool-deployment.com"
|
||||
srv, err := aibridgedserver.NewServer(t.Context(), db, logger, accessURL, tc.externalAuthConfigs, tc.experiments)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, srv)
|
||||
|
||||
resp, err := srv.GetMCPServerConfigs(t.Context(), &proto.GetMCPServerConfigsRequest{})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, resp)
|
||||
|
||||
if tc.expectCoderMCP {
|
||||
coderConfig := resp.CoderMcpConfig
|
||||
require.NotNil(t, coderConfig)
|
||||
require.Equal(t, "coder", coderConfig.GetId())
|
||||
expectedURL, err := url.JoinPath(accessURL, codermcp.MCPEndpoint)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, expectedURL, coderConfig.GetUrl())
|
||||
require.Empty(t, coderConfig.GetToolAllowRegex())
|
||||
require.Empty(t, coderConfig.GetToolDenyRegex())
|
||||
} else {
|
||||
require.Empty(t, resp.GetCoderMcpConfig())
|
||||
}
|
||||
|
||||
if tc.expectedExternalMCP {
|
||||
require.Len(t, resp.GetExternalAuthMcpConfigs(), 1)
|
||||
} else {
|
||||
require.Empty(t, resp.GetExternalAuthMcpConfigs())
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetMCPServerAccessTokensBatch(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctrl := gomock.NewController(t)
|
||||
db := dbmock.NewMockStore(ctrl)
|
||||
logger := testutil.Logger(t)
|
||||
|
||||
// Given: 2 external auth configured with MCP and 1 without.
|
||||
srv, err := aibridgedserver.NewServer(t.Context(), db, logger, "/", []*externalauth.Config{
|
||||
{
|
||||
ID: "1",
|
||||
MCPURL: "1.com/mcp",
|
||||
},
|
||||
{
|
||||
ID: "2",
|
||||
MCPURL: "2.com/mcp",
|
||||
},
|
||||
{
|
||||
ID: "3",
|
||||
},
|
||||
}, requiredExperiments)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, srv)
|
||||
|
||||
// When: requesting all external auth links, return all.
|
||||
db.EXPECT().GetExternalAuthLinksByUserID(gomock.Any(), gomock.Any()).MinTimes(1).DoAndReturn(func(ctx context.Context, userID uuid.UUID) ([]database.ExternalAuthLink, error) {
|
||||
return []database.ExternalAuthLink{
|
||||
{
|
||||
UserID: userID,
|
||||
ProviderID: "1",
|
||||
OAuthAccessToken: "1-token",
|
||||
},
|
||||
{
|
||||
UserID: userID,
|
||||
ProviderID: "2",
|
||||
OAuthAccessToken: "2-token",
|
||||
OAuthExpiry: dbtime.Now().Add(-time.Minute), // This token is expired and should not be returned.
|
||||
},
|
||||
{
|
||||
UserID: userID,
|
||||
ProviderID: "3",
|
||||
OAuthAccessToken: "3-token",
|
||||
},
|
||||
}, nil
|
||||
})
|
||||
|
||||
// When: accessing the MCP server access tokens, only the 2 with MCP configured should be returned, and the 1 without should
|
||||
// not fail the request but rather have an error returned specifically for that server.
|
||||
resp, err := srv.GetMCPServerAccessTokensBatch(t.Context(), &proto.GetMCPServerAccessTokensBatchRequest{
|
||||
UserId: uuid.NewString(),
|
||||
McpServerConfigIds: []string{"1", "1", "2", "3"}, // Duplicates must be tolerated.
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Then: 2 MCP servers are eligible but only 1 will return a valid token as the other expired.
|
||||
require.Len(t, resp.GetAccessTokens(), 1)
|
||||
require.Equal(t, "1-token", resp.GetAccessTokens()["1"])
|
||||
require.Len(t, resp.GetErrors(), 2)
|
||||
require.Contains(t, resp.GetErrors()["2"], aibridgedserver.ErrExpiredOrInvalidOAuthToken.Error())
|
||||
require.Contains(t, resp.GetErrors()["3"], aibridgedserver.ErrNoMCPConfigFound.Error())
|
||||
}
|
||||
|
||||
func TestRecordInterception(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var (
|
||||
metadataProto = map[string]*anypb.Any{
|
||||
"key": mustMarshalAny(t, &structpb.Value{Kind: &structpb.Value_StringValue{StringValue: "value"}}),
|
||||
}
|
||||
metadataJSON = `{"key":"value"}`
|
||||
)
|
||||
|
||||
testRecordMethod(t,
|
||||
func(srv *aibridgedserver.Server, ctx context.Context, req *proto.RecordInterceptionRequest) (*proto.RecordInterceptionResponse, error) {
|
||||
return srv.RecordInterception(ctx, req)
|
||||
},
|
||||
[]testRecordMethodCase[*proto.RecordInterceptionRequest]{
|
||||
{
|
||||
name: "valid interception",
|
||||
request: &proto.RecordInterceptionRequest{
|
||||
Id: uuid.NewString(),
|
||||
InitiatorId: uuid.NewString(),
|
||||
Provider: "anthropic",
|
||||
Model: "claude-4-opus",
|
||||
Metadata: metadataProto,
|
||||
StartedAt: timestamppb.Now(),
|
||||
},
|
||||
setupMocks: func(t *testing.T, db *dbmock.MockStore, req *proto.RecordInterceptionRequest) {
|
||||
interceptionID, err := uuid.Parse(req.GetId())
|
||||
assert.NoError(t, err, "parse interception UUID")
|
||||
initiatorID, err := uuid.Parse(req.GetInitiatorId())
|
||||
assert.NoError(t, err, "parse interception initiator UUID")
|
||||
|
||||
db.EXPECT().InsertAIBridgeInterception(gomock.Any(), database.InsertAIBridgeInterceptionParams{
|
||||
ID: interceptionID,
|
||||
InitiatorID: initiatorID,
|
||||
Provider: req.GetProvider(),
|
||||
Model: req.GetModel(),
|
||||
Metadata: json.RawMessage(metadataJSON),
|
||||
StartedAt: req.StartedAt.AsTime().UTC(),
|
||||
}).Return(database.AIBridgeInterception{
|
||||
ID: interceptionID,
|
||||
InitiatorID: initiatorID,
|
||||
Provider: req.GetProvider(),
|
||||
Model: req.GetModel(),
|
||||
StartedAt: req.StartedAt.AsTime().UTC(),
|
||||
}, nil)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "invalid interception ID",
|
||||
request: &proto.RecordInterceptionRequest{
|
||||
Id: "not-a-uuid",
|
||||
InitiatorId: uuid.NewString(),
|
||||
Provider: "anthropic",
|
||||
Model: "claude-4-opus",
|
||||
StartedAt: timestamppb.Now(),
|
||||
},
|
||||
expectedErr: "invalid interception ID",
|
||||
},
|
||||
{
|
||||
name: "invalid initiator ID",
|
||||
request: &proto.RecordInterceptionRequest{
|
||||
Id: uuid.NewString(),
|
||||
InitiatorId: "not-a-uuid",
|
||||
Provider: "anthropic",
|
||||
Model: "claude-4-opus",
|
||||
StartedAt: timestamppb.Now(),
|
||||
},
|
||||
expectedErr: "invalid initiator ID",
|
||||
},
|
||||
{
|
||||
name: "database error",
|
||||
request: &proto.RecordInterceptionRequest{
|
||||
Id: uuid.NewString(),
|
||||
InitiatorId: uuid.NewString(),
|
||||
Provider: "anthropic",
|
||||
Model: "claude-4-opus",
|
||||
StartedAt: timestamppb.Now(),
|
||||
},
|
||||
setupMocks: func(t *testing.T, db *dbmock.MockStore, req *proto.RecordInterceptionRequest) {
|
||||
db.EXPECT().InsertAIBridgeInterception(gomock.Any(), gomock.Any()).Return(database.AIBridgeInterception{}, sql.ErrConnDone)
|
||||
},
|
||||
expectedErr: "start interception",
|
||||
},
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
func TestRecordTokenUsage(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var (
|
||||
metadataProto = map[string]*anypb.Any{
|
||||
"key": mustMarshalAny(t, &structpb.Value{Kind: &structpb.Value_StringValue{StringValue: "value"}}),
|
||||
}
|
||||
metadataJSON = `{"key":"value"}`
|
||||
)
|
||||
|
||||
testRecordMethod(t,
|
||||
func(srv *aibridgedserver.Server, ctx context.Context, req *proto.RecordTokenUsageRequest) (*proto.RecordTokenUsageResponse, error) {
|
||||
return srv.RecordTokenUsage(ctx, req)
|
||||
},
|
||||
[]testRecordMethodCase[*proto.RecordTokenUsageRequest]{
|
||||
{
|
||||
name: "valid token usage",
|
||||
request: &proto.RecordTokenUsageRequest{
|
||||
InterceptionId: uuid.NewString(),
|
||||
MsgId: "msg_123",
|
||||
InputTokens: 100,
|
||||
OutputTokens: 200,
|
||||
Metadata: metadataProto,
|
||||
CreatedAt: timestamppb.Now(),
|
||||
},
|
||||
setupMocks: func(t *testing.T, db *dbmock.MockStore, req *proto.RecordTokenUsageRequest) {
|
||||
interceptionID, err := uuid.Parse(req.GetInterceptionId())
|
||||
assert.NoError(t, err, "parse interception UUID")
|
||||
|
||||
db.EXPECT().InsertAIBridgeTokenUsage(gomock.Any(), gomock.Cond(func(p database.InsertAIBridgeTokenUsageParams) bool {
|
||||
if !assert.NotEqual(t, uuid.Nil, p.ID, "ID") ||
|
||||
!assert.Equal(t, interceptionID, p.InterceptionID, "interception ID") ||
|
||||
!assert.Equal(t, req.GetMsgId(), p.ProviderResponseID, "provider response ID") ||
|
||||
!assert.Equal(t, req.GetInputTokens(), p.InputTokens, "input tokens") ||
|
||||
!assert.Equal(t, req.GetOutputTokens(), p.OutputTokens, "output tokens") ||
|
||||
!assert.JSONEq(t, metadataJSON, string(p.Metadata), "metadata") ||
|
||||
!assert.WithinDuration(t, req.GetCreatedAt().AsTime(), p.CreatedAt, time.Second, "created at") {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
})).Return(nil)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "invalid interception ID",
|
||||
request: &proto.RecordTokenUsageRequest{
|
||||
InterceptionId: "not-a-uuid",
|
||||
MsgId: "msg_123",
|
||||
InputTokens: 100,
|
||||
OutputTokens: 200,
|
||||
CreatedAt: timestamppb.Now(),
|
||||
},
|
||||
expectedErr: "failed to parse interception_id",
|
||||
},
|
||||
{
|
||||
name: "database error",
|
||||
request: &proto.RecordTokenUsageRequest{
|
||||
InterceptionId: uuid.NewString(),
|
||||
MsgId: "msg_123",
|
||||
InputTokens: 100,
|
||||
OutputTokens: 200,
|
||||
CreatedAt: timestamppb.Now(),
|
||||
},
|
||||
setupMocks: func(t *testing.T, db *dbmock.MockStore, req *proto.RecordTokenUsageRequest) {
|
||||
db.EXPECT().InsertAIBridgeTokenUsage(gomock.Any(), gomock.Any()).Return(sql.ErrConnDone)
|
||||
},
|
||||
expectedErr: "insert token usage",
|
||||
},
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
func TestRecordPromptUsage(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var (
|
||||
metadataProto = map[string]*anypb.Any{
|
||||
"key": mustMarshalAny(t, &structpb.Value{Kind: &structpb.Value_StringValue{StringValue: "value"}}),
|
||||
}
|
||||
metadataJSON = `{"key":"value"}`
|
||||
)
|
||||
|
||||
testRecordMethod(t,
|
||||
func(srv *aibridgedserver.Server, ctx context.Context, req *proto.RecordPromptUsageRequest) (*proto.RecordPromptUsageResponse, error) {
|
||||
return srv.RecordPromptUsage(ctx, req)
|
||||
},
|
||||
[]testRecordMethodCase[*proto.RecordPromptUsageRequest]{
|
||||
{
|
||||
name: "valid prompt usage",
|
||||
request: &proto.RecordPromptUsageRequest{
|
||||
InterceptionId: uuid.NewString(),
|
||||
MsgId: "msg_123",
|
||||
Prompt: "yo",
|
||||
Metadata: metadataProto,
|
||||
CreatedAt: timestamppb.Now(),
|
||||
},
|
||||
setupMocks: func(t *testing.T, db *dbmock.MockStore, req *proto.RecordPromptUsageRequest) {
|
||||
interceptionID, err := uuid.Parse(req.GetInterceptionId())
|
||||
assert.NoError(t, err, "parse interception UUID")
|
||||
|
||||
db.EXPECT().InsertAIBridgeUserPrompt(gomock.Any(), gomock.Cond(func(p database.InsertAIBridgeUserPromptParams) bool {
|
||||
if !assert.NotEqual(t, uuid.Nil, p.ID, "ID") ||
|
||||
!assert.Equal(t, interceptionID, p.InterceptionID, "interception ID") ||
|
||||
!assert.Equal(t, req.GetMsgId(), p.ProviderResponseID, "provider response ID") ||
|
||||
!assert.Equal(t, req.GetPrompt(), p.Prompt, "prompt") ||
|
||||
!assert.JSONEq(t, metadataJSON, string(p.Metadata), "metadata") ||
|
||||
!assert.WithinDuration(t, req.GetCreatedAt().AsTime(), p.CreatedAt, time.Second, "created at") {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
})).Return(nil)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "invalid interception ID",
|
||||
request: &proto.RecordPromptUsageRequest{
|
||||
InterceptionId: "not-a-uuid",
|
||||
MsgId: "msg_123",
|
||||
Prompt: "yo",
|
||||
CreatedAt: timestamppb.Now(),
|
||||
},
|
||||
expectedErr: "failed to parse interception_id",
|
||||
},
|
||||
{
|
||||
name: "database error",
|
||||
request: &proto.RecordPromptUsageRequest{
|
||||
InterceptionId: uuid.NewString(),
|
||||
MsgId: "msg_123",
|
||||
Prompt: "yo",
|
||||
CreatedAt: timestamppb.Now(),
|
||||
},
|
||||
setupMocks: func(t *testing.T, db *dbmock.MockStore, req *proto.RecordPromptUsageRequest) {
|
||||
db.EXPECT().InsertAIBridgeUserPrompt(gomock.Any(), gomock.Any()).Return(sql.ErrConnDone)
|
||||
},
|
||||
expectedErr: "insert user prompt",
|
||||
},
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
func TestRecordToolUsage(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var (
|
||||
metadataProto = map[string]*anypb.Any{
|
||||
"key": mustMarshalAny(t, &structpb.Value{Kind: &structpb.Value_NumberValue{NumberValue: 123.45}}),
|
||||
}
|
||||
metadataJSON = `{"key":123.45}`
|
||||
)
|
||||
|
||||
testRecordMethod(t,
|
||||
func(srv *aibridgedserver.Server, ctx context.Context, req *proto.RecordToolUsageRequest) (*proto.RecordToolUsageResponse, error) {
|
||||
return srv.RecordToolUsage(ctx, req)
|
||||
},
|
||||
[]testRecordMethodCase[*proto.RecordToolUsageRequest]{
|
||||
{
|
||||
name: "valid tool usage with all fields",
|
||||
request: &proto.RecordToolUsageRequest{
|
||||
InterceptionId: uuid.NewString(),
|
||||
MsgId: "msg_123",
|
||||
ServerUrl: strPtr("https://api.example.com"),
|
||||
Tool: "read_file",
|
||||
Input: `{"path": "/etc/hosts"}`,
|
||||
Injected: false,
|
||||
InvocationError: strPtr("permission denied"),
|
||||
Metadata: metadataProto,
|
||||
CreatedAt: timestamppb.Now(),
|
||||
},
|
||||
setupMocks: func(t *testing.T, db *dbmock.MockStore, req *proto.RecordToolUsageRequest) {
|
||||
interceptionID, err := uuid.Parse(req.GetInterceptionId())
|
||||
assert.NoError(t, err, "parse interception UUID")
|
||||
|
||||
dbServerURL := sql.NullString{}
|
||||
if req.ServerUrl != nil {
|
||||
dbServerURL.String = *req.ServerUrl
|
||||
dbServerURL.Valid = true
|
||||
}
|
||||
|
||||
dbInvocationError := sql.NullString{}
|
||||
if req.InvocationError != nil {
|
||||
dbInvocationError.String = *req.InvocationError
|
||||
dbInvocationError.Valid = true
|
||||
}
|
||||
|
||||
db.EXPECT().InsertAIBridgeToolUsage(gomock.Any(), gomock.Cond(func(p database.InsertAIBridgeToolUsageParams) bool {
|
||||
if !assert.NotEqual(t, uuid.Nil, p.ID, "ID") ||
|
||||
!assert.Equal(t, interceptionID, p.InterceptionID, "interception ID") ||
|
||||
!assert.Equal(t, req.GetMsgId(), p.ProviderResponseID, "provider response ID") ||
|
||||
!assert.Equal(t, req.GetTool(), p.Tool, "tool") ||
|
||||
!assert.Equal(t, dbServerURL, p.ServerUrl, "server URL") ||
|
||||
!assert.Equal(t, req.GetInput(), p.Input, "input") ||
|
||||
!assert.Equal(t, req.GetInjected(), p.Injected, "injected") ||
|
||||
!assert.Equal(t, dbInvocationError, p.InvocationError, "invocation error") ||
|
||||
!assert.JSONEq(t, metadataJSON, string(p.Metadata), "metadata") ||
|
||||
!assert.WithinDuration(t, req.GetCreatedAt().AsTime(), p.CreatedAt, time.Second, "created at") {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
})).Return(nil)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "invalid interception ID",
|
||||
request: &proto.RecordToolUsageRequest{
|
||||
InterceptionId: "not-a-uuid",
|
||||
MsgId: "msg_123",
|
||||
Tool: "read_file",
|
||||
Input: `{"path": "/etc/hosts"}`,
|
||||
CreatedAt: timestamppb.Now(),
|
||||
},
|
||||
expectedErr: "failed to parse interception_id",
|
||||
},
|
||||
{
|
||||
name: "database error",
|
||||
request: &proto.RecordToolUsageRequest{
|
||||
InterceptionId: uuid.NewString(),
|
||||
MsgId: "msg_123",
|
||||
Tool: "read_file",
|
||||
Input: `{"path": "/etc/hosts"}`,
|
||||
CreatedAt: timestamppb.Now(),
|
||||
},
|
||||
setupMocks: func(t *testing.T, db *dbmock.MockStore, req *proto.RecordToolUsageRequest) {
|
||||
db.EXPECT().InsertAIBridgeToolUsage(gomock.Any(), gomock.Any()).Return(sql.ErrConnDone)
|
||||
},
|
||||
expectedErr: "insert tool usage",
|
||||
},
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
type testRecordMethodCase[Req any] struct {
|
||||
name string
|
||||
request Req
|
||||
// setupMocks is called with the mock store and the above request.
|
||||
setupMocks func(t *testing.T, db *dbmock.MockStore, req Req)
|
||||
expectedErr string
|
||||
}
|
||||
|
||||
// testRecordMethod is a helper that abstracts the common testing pattern for all Record* methods.
|
||||
func testRecordMethod[Req any, Resp any](
|
||||
t *testing.T,
|
||||
callMethod func(srv *aibridgedserver.Server, ctx context.Context, req Req) (Resp, error),
|
||||
cases []testRecordMethodCase[Req],
|
||||
) {
|
||||
t.Helper()
|
||||
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctrl := gomock.NewController(t)
|
||||
db := dbmock.NewMockStore(ctrl)
|
||||
logger := testutil.Logger(t)
|
||||
|
||||
if tc.setupMocks != nil {
|
||||
tc.setupMocks(t, db, tc.request)
|
||||
}
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
srv, err := aibridgedserver.NewServer(ctx, db, logger, "/", nil, requiredExperiments)
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := callMethod(srv, ctx, tc.request)
|
||||
if tc.expectedErr != "" {
|
||||
require.Error(t, err, "Expected error for test case: %s", tc.name)
|
||||
require.Contains(t, err.Error(), tc.expectedErr)
|
||||
} else {
|
||||
require.NoError(t, err, "Unexpected error for test case: %s", tc.name)
|
||||
require.NotNil(t, resp)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Helper functions.
|
||||
func mustMarshalAny(t *testing.T, msg protobufproto.Message) *anypb.Any {
|
||||
t.Helper()
|
||||
v, err := anypb.New(msg)
|
||||
require.NoError(t, err)
|
||||
return v
|
||||
}
|
||||
|
||||
func strPtr(s string) *string {
|
||||
return &s
|
||||
}
|
||||
Reference in New Issue
Block a user