diff --git a/.ai/prompts/tests.md b/.ai/prompts/tests.md new file mode 100644 index 000000000..461e689b3 --- /dev/null +++ b/.ai/prompts/tests.md @@ -0,0 +1,68 @@ + +# Tests + +- Prefer `go test ` or `-run ` over `go test ./...` (slow). +- Use table-driven tests covering happy path, failure, and edge cases. +- Skip trivial getters/setters unless they contain non-trivial logic. +- Use `testify/assert` with `assert.*(t, *)` or `require.*(t, *)` directly, not `assert.New(t)`. +- Use `testify/suite` for related test groups, and use `s.*(*, *)` when asserting. +- Use the testify `EXPECT` method for mocks; avoid `mock.Anything`. +- Use `assert.AnError` if needed, and `assert.Equal` for error assertions. +- Assert full maps/structs/slices/arrays, not individual fields. +- Prefer direct value assertions over `mock.MatchedBy`; use it only for dynamically-generated args. +- Every mock must use `.Once()` or `.Times()`, only use `.Maybe()` when necessary; avoid no-op expectations. +- Name tests `Test_[Optional]`; use table style or sub-tests for multiple cases. +- Don't use `assert.*` with `if` statements; use `assert.*` directly for clarity and better failure messages. +- Use `t.Run()` for sub-tests when testing multiple cases for the same function, and use table-driven tests for multiple cases with similar setup/assertions. Avoid writing separate test functions for each case when they share common logic. +- The basic table-driven test pattern is: + +```go +import ( + [system packages] + + [third-party packages] + + [internal packages] +) + +func TestFunction(t *testing.T) { + // The name should start with `mock` to indicate it's a mocked function. + var ( + ctx context.Context + mockFunc *mocks.MockedInterface + ) + + beforeEach := func() { + mockFunc = mocks.NewMockedInterface(t) + } + + tests := []struct { + name string + input any + setup func() + expect any + expectError error + }{ + { + name: "should do something", + input: someInput, + setup: func() { + mockFunc.EXPECT().SomeMethod(someArgs).Return(someResult, nil).Once() + }, + expect: someResult, + expectError: nil, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + beforeEach() + tt.setup() + + result, err := FunctionUnderTest(tt.input) + assert.Equal(t, tt.expect, result) + assert.Equal(t, tt.expectError, err) + }) + } +} +``` diff --git a/AGENTS.md b/AGENTS.md index 10e5751d2..9224aaf25 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -40,10 +40,12 @@ golangci-lint run # lint - Cache/Session/Queue: `goravel/redis` - Storage: `goravel/s3`, `oss`, `cos`, `minio` -## AI Agent Code Rule +## Code Rules -- Should use `any` instead of `interface{}`. -- Avoid adding `mock.Anything` when writing test cases. -- Use the testify `EXPECT` method when writing test cases. -- Don't modify the files in the `mocks` directory, run the `go tool mockery` command to regenerate mocks instead if needed. -- Don't run `go test ./...` if unnecessary, the command is a bit slow, run `go test` with the specific package or test function instead. \ No newline at end of file +- Use `any` instead of `interface{}`. +- Never edit `mocks/` directly; run `go tool mockery` to regenerate. +- Follow standard Go formatting/naming; add comments where logic isn't self-evident. Go version is in go.mod. + +## Tests + +When writing tests, use the rules in `.ai/prompts/tests.md` for guidance. diff --git a/ai/application.go b/ai/application.go index 88709d5ee..5c1a2f703 100644 --- a/ai/application.go +++ b/ai/application.go @@ -1,20 +1,52 @@ package ai import ( + "context" + contractsai "github.com/goravel/framework/contracts/ai" - "github.com/goravel/framework/contracts/config" ) var _ contractsai.AI = (*Application)(nil) -// Application is the AI manager implementation. type Application struct { + ctx context.Context + config contractsai.Config + resolver *ProviderResolver } -func NewApplication(config config.Config) *Application { - return &Application{} +func NewApplication(ctx context.Context, config contractsai.Config) *Application { + return &Application{ + ctx: ctx, + config: config, + resolver: NewProviderResolver(config), + } } func (r *Application) Agent(agent contractsai.Agent, options ...contractsai.Option) (contractsai.Conversation, error) { - return &conversation{}, nil + opts := make(map[string]any) + for _, option := range options { + option(opts) + } + + providerName, _ := opts[contractsai.OptionProvider].(string) + if providerName == "" { + providerName = r.config.Default + } + + provider, err := r.resolver.New(providerName) + if err != nil { + return nil, err + } + + model, _ := opts[contractsai.OptionModel].(string) + + return NewConversation(r.ctx, agent, provider, model), nil +} + +func (r *Application) WithContext(ctx context.Context) contractsai.AI { + return &Application{ + ctx: ctx, + config: r.config, + resolver: r.resolver, + } } diff --git a/ai/application_test.go b/ai/application_test.go new file mode 100644 index 000000000..d260661c1 --- /dev/null +++ b/ai/application_test.go @@ -0,0 +1,174 @@ +package ai + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + + contractsai "github.com/goravel/framework/contracts/ai" + mocksai "github.com/goravel/framework/mocks/ai" +) + +func TestApplication_Agent(t *testing.T) { + ctx := context.Background() + tests := []struct { + name string + promptInput string + options []contractsai.Option + setupConfig func(t *testing.T) (contractsai.Config, *mocksai.Provider) + expectedModel string + responseText string + expectResponse bool + promptErr error + expectPromptErr bool + }{ + { + name: "default provider", + promptInput: "ping", + setupConfig: func(t *testing.T) (contractsai.Config, *mocksai.Provider) { + provider := mocksai.NewProvider(t) + return contractsai.Config{ + Default: "default", + Providers: map[string]contractsai.ProviderConfig{ + "default": {Via: provider}, + }, + }, provider + }, + responseText: "ok", + expectResponse: true, + }, + { + name: "provider override", + promptInput: "override", + options: []contractsai.Option{WithProvider("alternative")}, + setupConfig: func(t *testing.T) (contractsai.Config, *mocksai.Provider) { + defaultProvider := mocksai.NewProvider(t) + alternativeProvider := mocksai.NewProvider(t) + return contractsai.Config{ + Default: "default", + Providers: map[string]contractsai.ProviderConfig{ + "default": {Via: defaultProvider}, + "alternative": {Via: alternativeProvider}, + }, + }, alternativeProvider + }, + responseText: "override", + expectResponse: true, + }, + { + name: "model option", + promptInput: "any", + options: []contractsai.Option{WithModel("custom-model")}, + setupConfig: func(t *testing.T) (contractsai.Config, *mocksai.Provider) { + provider := mocksai.NewProvider(t) + return contractsai.Config{ + Default: "default", + Providers: map[string]contractsai.ProviderConfig{ + "default": {Via: provider}, + }, + }, provider + }, + expectedModel: "custom-model", + responseText: "modelled", + expectResponse: true, + }, + { + name: "provider error", + promptInput: "fail", + setupConfig: func(t *testing.T) (contractsai.Config, *mocksai.Provider) { + provider := mocksai.NewProvider(t) + return contractsai.Config{ + Default: "default", + Providers: map[string]contractsai.ProviderConfig{ + "default": {Via: provider}, + }, + }, provider + }, + expectPromptErr: true, + promptErr: assert.AnError, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + config, provider := tt.setupConfig(t) + agent := mocksai.NewAgent(t) + agent.EXPECT().Messages().Return(nil).Once() + + app := NewApplication(ctx, config) + conv, err := app.Agent(agent, tt.options...) + assert.NoError(t, err) + + convImpl, ok := conv.(*conversation) + assert.True(t, ok) + + expectedPrompt := contractsai.AgentPrompt{ + Agent: convImpl, + Input: tt.promptInput, + Model: tt.expectedModel, + } + + var response *mocksai.Response + if tt.expectResponse { + response = mocksai.NewResponse(t) + response.EXPECT().Text().Return(tt.responseText).Once() + } + + provider.EXPECT(). + Prompt(ctx, expectedPrompt). + Return(response, tt.promptErr). + Once() + + resp, err := conv.Prompt(tt.promptInput) + if tt.expectPromptErr { + assert.Equal(t, tt.promptErr, err) + assert.Nil(t, resp) + return + } + assert.NoError(t, err) + assert.Equal(t, response, resp) + }) + } +} + +func TestApplication_Agent_ResolverError(t *testing.T) { + ctx := context.Background() + config := contractsai.Config{ + Default: "default", + Providers: map[string]contractsai.ProviderConfig{ + "default": { + Via: func() (contractsai.Provider, error) { + return nil, assert.AnError + }, + }, + }, + } + + app := NewApplication(ctx, config) + _, err := app.Agent(mocksai.NewAgent(t)) + assert.Equal(t, assert.AnError, err) +} + +type testCtxKey string + +func TestApplication_WithContext(t *testing.T) { + origCtx := context.WithValue(context.Background(), testCtxKey("orig"), true) + provider := mocksai.NewProvider(t) + config := contractsai.Config{ + Default: "default", + Providers: map[string]contractsai.ProviderConfig{ + "default": {Via: provider}, + }, + } + + app := NewApplication(origCtx, config) + newCtx := context.WithValue(context.Background(), testCtxKey("orig"), false) + aiWithCtx := app.WithContext(newCtx) + aiImpl, ok := aiWithCtx.(*Application) + assert.True(t, ok) + + assert.Same(t, newCtx, aiImpl.ctx) + assert.Same(t, app.resolver, aiImpl.resolver) + assert.Equal(t, app.config, aiImpl.config) +} diff --git a/ai/conversation.go b/ai/conversation.go index 754fba956..05ebe6784 100644 --- a/ai/conversation.go +++ b/ai/conversation.go @@ -2,20 +2,48 @@ package ai import ( "context" + "slices" contractsai "github.com/goravel/framework/contracts/ai" ) type conversation struct { + ctx context.Context + agent contractsai.Agent + messages []contractsai.Message + provider contractsai.Provider + model string } -func (r *conversation) Prompt(ctx context.Context, input string) (contractsai.Response, error) { - return nil, nil +func NewConversation(ctx context.Context, agent contractsai.Agent, provider contractsai.Provider, model string) *conversation { + return &conversation{ + ctx: ctx, + agent: agent, + messages: slices.Clone(agent.Messages()), + provider: provider, + model: model, + } } -func (r *conversation) Messages() []contractsai.Message { - return nil -} +func (r *conversation) Instructions() string { return r.agent.Instructions() } +func (r *conversation) Messages() []contractsai.Message { return r.messages } + +func (r *conversation) Prompt(input string) (contractsai.Response, error) { + resp, err := r.provider.Prompt(r.ctx, contractsai.AgentPrompt{ + Agent: r, + Input: input, + Model: r.model, + }) + if err != nil { + return nil, err + } -func (r *conversation) Reset() { + r.messages = append(r.messages, + contractsai.Message{Role: contractsai.RoleUser, Content: input}, + contractsai.Message{Role: contractsai.RoleAssistant, Content: resp.Text()}, + ) + + return resp, nil } + +func (r *conversation) Reset() { r.messages = slices.Clone(r.agent.Messages()) } diff --git a/ai/conversation_test.go b/ai/conversation_test.go new file mode 100644 index 000000000..f73fd70c7 --- /dev/null +++ b/ai/conversation_test.go @@ -0,0 +1,168 @@ +package ai + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + contractsai "github.com/goravel/framework/contracts/ai" + mocksai "github.com/goravel/framework/mocks/ai" +) + +func TestConversation_Prompt(t *testing.T) { + ctx := context.Background() + + var ( + mockProvider *mocksai.Provider + conv *conversation + ) + + beforeEach := func(initial []contractsai.Message, model string) { + mockAgent := mocksai.NewAgent(t) + mockAgent.EXPECT().Messages().Return(initial).Once() + mockProvider = mocksai.NewProvider(t) + conv = NewConversation(ctx, mockAgent, mockProvider, model) + } + + tests := []struct { + name string + initial []contractsai.Message + model string + input string + setup func() contractsai.Response + expectMessages []contractsai.Message + expectError error + }{ + { + name: "appends messages on success", + initial: []contractsai.Message{{Role: contractsai.RoleUser, Content: "system"}}, + model: "model-x", + input: "hello", + setup: func() contractsai.Response { + mockResponse := mocksai.NewResponse(t) + mockResponse.EXPECT().Text().Return("got it").Once() + mockProvider.EXPECT().Prompt(ctx, contractsai.AgentPrompt{Agent: conv, Input: "hello", Model: "model-x"}).Return(mockResponse, nil).Once() + return mockResponse + }, + expectMessages: []contractsai.Message{ + {Role: contractsai.RoleUser, Content: "system"}, + {Role: contractsai.RoleUser, Content: "hello"}, + {Role: contractsai.RoleAssistant, Content: "got it"}, + }, + }, + { + name: "does not append on error", + initial: []contractsai.Message{{Role: contractsai.RoleAssistant, Content: "init"}}, + model: "model-y", + input: "fail", + setup: func() contractsai.Response { + mockProvider.EXPECT().Prompt(ctx, contractsai.AgentPrompt{Agent: conv, Input: "fail", Model: "model-y"}).Return(nil, assert.AnError).Once() + return nil + }, + expectMessages: []contractsai.Message{ + {Role: contractsai.RoleAssistant, Content: "init"}, + }, + expectError: assert.AnError, + }, + { + name: "appends empty input and empty response", + initial: []contractsai.Message{}, + model: "model-empty", + input: "", + setup: func() contractsai.Response { + mockResponse := mocksai.NewResponse(t) + mockResponse.EXPECT().Text().Return("").Once() + mockProvider.EXPECT().Prompt(ctx, contractsai.AgentPrompt{Agent: conv, Input: "", Model: "model-empty"}).Return(mockResponse, nil).Once() + return mockResponse + }, + expectMessages: []contractsai.Message{ + {Role: contractsai.RoleUser, Content: ""}, + {Role: contractsai.RoleAssistant, Content: ""}, + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + beforeEach(tt.initial, tt.model) + expectResp := tt.setup() + + resp, err := conv.Prompt(tt.input) + assert.Equal(t, tt.expectError, err) + assert.Equal(t, expectResp, resp) + assert.Equal(t, tt.expectMessages, conv.Messages()) + }) + } +} + +func TestConversation_Reset(t *testing.T) { + ctx := context.Background() + + var ( + mockProvider *mocksai.Provider + conv *conversation + ) + + beforeEach := func(initial []contractsai.Message) { + mockAgent := mocksai.NewAgent(t) + mockAgent.EXPECT().Messages().Return(initial).Times(2) + mockProvider = mocksai.NewProvider(t) + conv = NewConversation(ctx, mockAgent, mockProvider, "model-z") + } + + tests := []struct { + name string + initial []contractsai.Message + input string + promptBefore bool + setup func() + expectBefore []contractsai.Message + expectAfter []contractsai.Message + }{ + { + name: "restores initial messages after prompt", + initial: []contractsai.Message{{Role: contractsai.RoleToolResult, Content: "keep"}}, + input: "append", + promptBefore: true, + setup: func() { + response := mocksai.NewResponse(t) + mockProvider.EXPECT().Prompt(ctx, contractsai.AgentPrompt{Agent: conv, Input: "append", Model: "model-z"}).Return(response, nil).Once() + response.EXPECT().Text().Return("done").Once() + }, + expectBefore: []contractsai.Message{ + {Role: contractsai.RoleToolResult, Content: "keep"}, + {Role: contractsai.RoleUser, Content: "append"}, + {Role: contractsai.RoleAssistant, Content: "done"}, + }, + expectAfter: []contractsai.Message{{Role: contractsai.RoleToolResult, Content: "keep"}}, + }, + { + name: "keeps same messages when reset without prompt", + initial: []contractsai.Message{{Role: contractsai.RoleAssistant, Content: "seed"}}, + promptBefore: false, + setup: func() {}, + expectBefore: []contractsai.Message{{Role: contractsai.RoleAssistant, Content: "seed"}}, + expectAfter: []contractsai.Message{{Role: contractsai.RoleAssistant, Content: "seed"}}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + beforeEach(tt.initial) + tt.setup() + + if tt.promptBefore { + _, err := conv.Prompt(tt.input) + require.NoError(t, err) + } + + assert.Equal(t, tt.expectBefore, conv.Messages()) + conv.Reset() + resetMessages := conv.Messages() + assert.Equal(t, tt.expectAfter, resetMessages) + assert.NotSame(t, &tt.initial[0], &resetMessages[0]) + }) + } +} diff --git a/ai/openai/provider.go b/ai/openai/provider.go new file mode 100644 index 000000000..85172c89c --- /dev/null +++ b/ai/openai/provider.go @@ -0,0 +1,80 @@ +package openai + +import ( + "context" + + goopenai "github.com/openai/openai-go" + "github.com/openai/openai-go/option" + + contractsai "github.com/goravel/framework/contracts/ai" + contractsconfig "github.com/goravel/framework/contracts/config" +) + +// The OpenAI provider will be moved into a separate package in the future. + +const DefaultTextModel = "gpt-5.4" + +type Provider struct { + client goopenai.Client + config contractsai.ProviderConfig +} + +func NewOpenAI(config contractsconfig.Config, provider string) (*Provider, error) { + var providerConfig contractsai.ProviderConfig + err := config.UnmarshalKey("ai.providers."+provider, &providerConfig) + if err != nil { + return nil, err + } + + opts := []option.RequestOption{option.WithAPIKey(providerConfig.Key)} + if providerConfig.Url != "" { + opts = append(opts, option.WithBaseURL(providerConfig.Url)) + } + if providerConfig.Models.Text.Default == "" { + providerConfig.Models.Text.Default = DefaultTextModel + } + + return &Provider{client: goopenai.NewClient(opts...), config: providerConfig}, nil +} + +func (r *Provider) Prompt(ctx context.Context, prompt contractsai.AgentPrompt) (contractsai.Response, error) { + model := r.config.Models.Text.Default + if prompt.Model != "" { + model = prompt.Model + } + + var messages []goopenai.ChatCompletionMessageParamUnion + if instructions := prompt.Agent.Instructions(); instructions != "" { + messages = append(messages, goopenai.SystemMessage(instructions)) + } + for _, m := range prompt.Agent.Messages() { + switch m.Role { + case contractsai.RoleUser: + messages = append(messages, goopenai.UserMessage(m.Content)) + case contractsai.RoleAssistant: + messages = append(messages, goopenai.AssistantMessage(m.Content)) + } + } + messages = append(messages, goopenai.UserMessage(prompt.Input)) + + completion, err := r.client.Chat.Completions.New(ctx, goopenai.ChatCompletionNewParams{ + Model: model, + Messages: messages, + }) + if err != nil { + return nil, err + } + + text := "" + if len(completion.Choices) > 0 { + text = completion.Choices[0].Message.Content + } + return &response{ + text: text, + usage: &usage{ + input: int(completion.Usage.PromptTokens), + output: int(completion.Usage.CompletionTokens), + total: int(completion.Usage.TotalTokens), + }, + }, nil +} diff --git a/ai/openai/provider_test.go b/ai/openai/provider_test.go new file mode 100644 index 000000000..4e9e025be --- /dev/null +++ b/ai/openai/provider_test.go @@ -0,0 +1,366 @@ +package openai + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" + + goopenai "github.com/openai/openai-go" + "github.com/openai/openai-go/option" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + contractsai "github.com/goravel/framework/contracts/ai" + "github.com/goravel/framework/errors" + mocksai "github.com/goravel/framework/mocks/ai" + mocksconfig "github.com/goravel/framework/mocks/config" +) + +type capturedRequest struct { + path string + authorization string + model string + messages []map[string]any +} + +type normalizedCapturedRequest struct { + path string + authorization string + model string + messages []normalizedMessage +} + +type normalizedMessage struct { + role string + content string +} + +func TestNewOpenAIUnmarshalError(t *testing.T) { + var mockConfig *mocksconfig.Config + + beforeEach := func() { + mockConfig = mocksconfig.NewConfig(t) + } + + tests := []struct { + name string + setup func() + expectConfig *contractsai.ProviderConfig + expectErr error + }{ + { + name: "returns unmarshal error", + setup: func() { + mockConfig.EXPECT().UnmarshalKey("ai.providers.openai", new(contractsai.ProviderConfig)).Return(assert.AnError).Once() + }, + expectErr: assert.AnError, + }, + { + name: "sets default text model", + setup: func() { + mockConfig.EXPECT().UnmarshalKey("ai.providers.openai", new(contractsai.ProviderConfig)).RunAndReturn(func(_ string, rawVal any) error { + cfg := rawVal.(*contractsai.ProviderConfig) + cfg.Key = "test-key" + cfg.Url = "http://localhost:1234" + return nil + }).Once() + }, + expectConfig: func() *contractsai.ProviderConfig { + cfg := contractsai.ProviderConfig{Key: "test-key", Url: "http://localhost:1234"} + cfg.Models.Text.Default = DefaultTextModel + return &cfg + }(), + }, + { + name: "keeps configured default model", + setup: func() { + mockConfig.EXPECT().UnmarshalKey("ai.providers.openai", new(contractsai.ProviderConfig)).RunAndReturn(func(_ string, rawVal any) error { + cfg := rawVal.(*contractsai.ProviderConfig) + cfg.Key = "test-key" + cfg.Models.Text.Default = "gpt-custom" + return nil + }).Once() + }, + expectConfig: func() *contractsai.ProviderConfig { + cfg := contractsai.ProviderConfig{Key: "test-key"} + cfg.Models.Text.Default = "gpt-custom" + return &cfg + }(), + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + beforeEach() + tt.setup() + + provider, err := NewOpenAI(mockConfig, "openai") + + assert.Equal(t, tt.expectErr, err) + if tt.expectErr != nil { + assert.Nil(t, provider) + return + } + require.NotNil(t, provider) + assert.Equal(t, *tt.expectConfig, provider.config) + }) + } +} + +func TestProviderPrompt(t *testing.T) { + type usageCheck struct { + input int + output int + total int + } + + var mockAgent *mocksai.Agent + + beforeEach := func() { + mockAgent = mocksai.NewAgent(t) + } + + tests := []struct { + name string + status int + body string + setup func() + modelOverride string + input string + expectText string + expectUsage usageCheck + expectErr bool + expectErrMessage string + expectRequest normalizedCapturedRequest + }{ + { + name: "builds messages with default model", + status: http.StatusOK, + body: `{"id":"cmpl_1","object":"chat.completion","created":1,"model":"gpt-test","choices":[{"index":0,"finish_reason":"stop","message":{"role":"assistant","content":"assistant reply","refusal":""}}],"usage":{"prompt_tokens":11,"completion_tokens":7,"total_tokens":18}}`, + setup: func() { + mockAgent.EXPECT().Instructions().Return("system rule").Once() + mockAgent.EXPECT().Messages().Return([]contractsai.Message{ + {Role: contractsai.RoleUser, Content: "history user"}, + {Role: contractsai.RoleAssistant, Content: "history assistant"}, + }).Once() + }, + input: "new input", + expectText: "assistant reply", + expectUsage: usageCheck{input: 11, output: 7, total: 18}, + expectRequest: normalizedCapturedRequest{ + path: "/chat/completions", + authorization: "Bearer test-key", + model: "gpt-default", + messages: []normalizedMessage{ + {role: "system", content: "system rule"}, + {role: "user", content: "history user"}, + {role: "assistant", content: "history assistant"}, + {role: "user", content: "new input"}, + }, + }, + }, + { + name: "uses prompt model override", + status: http.StatusOK, + body: `{"id":"cmpl_2","object":"chat.completion","created":1,"model":"gpt-test","choices":[{"index":0,"finish_reason":"stop","message":{"role":"assistant","content":"ok","refusal":""}}],"usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}`, + setup: func() { + mockAgent.EXPECT().Instructions().Return("").Once() + mockAgent.EXPECT().Messages().Return(nil).Once() + }, + modelOverride: "gpt-override", + input: "hello", + expectText: "ok", + expectUsage: usageCheck{input: 1, output: 1, total: 2}, + expectRequest: normalizedCapturedRequest{ + path: "/chat/completions", + authorization: "Bearer test-key", + model: "gpt-override", + messages: []normalizedMessage{ + {role: "user", content: "hello"}, + }, + }, + }, + { + name: "returns error when API fails", + status: http.StatusInternalServerError, + body: `{"error":{"message":"boom","type":"server_error"}}`, + setup: func() { + mockAgent.EXPECT().Instructions().Return("").Once() + mockAgent.EXPECT().Messages().Return(nil).Once() + }, + input: "hello", + expectErr: true, + expectErrMessage: "boom", + expectRequest: normalizedCapturedRequest{ + path: "/chat/completions", + authorization: "Bearer test-key", + model: "gpt-default", + messages: []normalizedMessage{ + {role: "user", content: "hello"}, + }, + }, + }, + { + name: "handles empty choices", + status: http.StatusOK, + body: `{"id":"cmpl_3","object":"chat.completion","created":1,"model":"gpt-test","choices":[],"usage":{"prompt_tokens":3,"completion_tokens":0,"total_tokens":3}}`, + setup: func() { + mockAgent.EXPECT().Instructions().Return("").Once() + mockAgent.EXPECT().Messages().Return(nil).Once() + }, + input: "hello", + expectText: "", + expectUsage: usageCheck{input: 3, output: 0, total: 3}, + expectRequest: normalizedCapturedRequest{ + path: "/chat/completions", + authorization: "Bearer test-key", + model: "gpt-default", + messages: []normalizedMessage{ + {role: "user", content: "hello"}, + }, + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + beforeEach() + + captured := make(chan capturedRequest, 1) + server := newChatServer(t, tt.status, tt.body, captured) + t.Cleanup(server.Close) + + provider := &Provider{ + client: goopenai.NewClient(option.WithBaseURL(server.URL), option.WithAPIKey("test-key")), + config: contractsai.ProviderConfig{}, + } + provider.config.Models.Text.Default = "gpt-default" + + tt.setup() + + prompt := contractsai.AgentPrompt{Agent: mockAgent, Input: tt.input} + if tt.modelOverride != "" { + prompt.Model = tt.modelOverride + } + + resp, err := provider.Prompt(context.Background(), prompt) + + if tt.expectErr { + assert.Nil(t, resp) + require.Error(t, err) + + var apiErr *goopenai.Error + require.ErrorAs(t, err, &apiErr) + assert.Equal(t, tt.expectErrMessage, apiErr.Message) + assert.ErrorContains(t, err, tt.expectErrMessage) + assert.Equal(t, tt.status, apiErr.StatusCode) + + req, ok := readCapturedRequest(t, captured) + require.True(t, ok, "expected request payload") + assert.Equal(t, tt.expectRequest, normalizeCapturedRequest(req)) + return + } + + require.NoError(t, err) + require.NotNil(t, resp) + assert.Equal(t, tt.expectText, resp.Text()) + require.NotNil(t, resp.Usage()) + assert.Equal(t, tt.expectUsage, usageCheck{ + input: resp.Usage().Input(), + output: resp.Usage().Output(), + total: resp.Usage().Total(), + }) + + req, ok := readCapturedRequest(t, captured) + require.True(t, ok, "expected request payload") + assert.Equal(t, tt.expectRequest, normalizeCapturedRequest(req)) + }) + } +} + +func newChatServer(t *testing.T, status int, body string, captured chan<- capturedRequest) *httptest.Server { + t.Helper() + + handler := func(w http.ResponseWriter, r *http.Request) { + defer errors.Ignore(r.Body.Close) + + var payload struct { + Model string `json:"model"` + Messages []map[string]any `json:"messages"` + } + if err := json.NewDecoder(r.Body).Decode(&payload); err == nil { + select { + case captured <- capturedRequest{ + path: r.URL.Path, + authorization: r.Header.Get("Authorization"), + model: payload.Model, + messages: payload.Messages, + }: + default: + } + } + + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(status) + _, _ = w.Write([]byte(body)) + } + + mux := http.NewServeMux() + mux.HandleFunc("/chat/completions", handler) + mux.HandleFunc("/v1/chat/completions", handler) + + return httptest.NewServer(mux) +} + +func messageText(content any) string { + switch val := content.(type) { + case string: + return val + case []any: + parts := make([]string, 0, len(val)) + for _, item := range val { + part, ok := item.(map[string]any) + if !ok { + continue + } + text, _ := part["text"].(string) + if text != "" { + parts = append(parts, text) + } + } + return strings.Join(parts, "") + default: + return "" + } +} + +func readCapturedRequest(t *testing.T, captured <-chan capturedRequest) (capturedRequest, bool) { + t.Helper() + select { + case req := <-captured: + return req, true + default: + return capturedRequest{}, false + } +} + +func normalizeCapturedRequest(req capturedRequest) normalizedCapturedRequest { + messages := make([]normalizedMessage, 0, len(req.messages)) + for _, message := range req.messages { + role, _ := message["role"].(string) + messages = append(messages, normalizedMessage{ + role: role, + content: messageText(message["content"]), + }) + } + + return normalizedCapturedRequest{ + path: req.path, + authorization: req.authorization, + model: req.model, + messages: messages, + } +} diff --git a/ai/openai/response.go b/ai/openai/response.go new file mode 100644 index 000000000..4d77a49c4 --- /dev/null +++ b/ai/openai/response.go @@ -0,0 +1,17 @@ +package openai + +import contractsai "github.com/goravel/framework/contracts/ai" + +type response struct { + text string + usage *usage +} + +func (r *response) Text() string { return r.text } +func (r *response) Usage() contractsai.Usage { return r.usage } + +type usage struct{ input, output, total int } + +func (r *usage) Input() int { return r.input } +func (r *usage) Output() int { return r.output } +func (r *usage) Total() int { return r.total } diff --git a/ai/option.go b/ai/option.go index 5f33565b0..c407c18d8 100644 --- a/ai/option.go +++ b/ai/option.go @@ -1,10 +1,6 @@ package ai -import ( - "time" - - contractsai "github.com/goravel/framework/contracts/ai" -) +import contractsai "github.com/goravel/framework/contracts/ai" func WithProvider(provider string) contractsai.Option { return func(options map[string]any) { @@ -17,9 +13,3 @@ func WithModel(model string) contractsai.Option { options[contractsai.OptionModel] = model } } - -func WithTimeout(timeout time.Duration) contractsai.Option { - return func(options map[string]any) { - options[contractsai.OptionTimeout] = timeout - } -} diff --git a/ai/option_test.go b/ai/option_test.go new file mode 100644 index 000000000..6a33e4a72 --- /dev/null +++ b/ai/option_test.go @@ -0,0 +1,113 @@ +package ai + +import ( + "testing" + + "github.com/stretchr/testify/assert" + + contractsai "github.com/goravel/framework/contracts/ai" +) + +func TestWithProvider(t *testing.T) { + tests := []struct { + name string + initial map[string]any + args []string + expected map[string]any + nilMap bool + }{ + { + name: "sets provider while preserving existing keys", + initial: map[string]any{"existing-key": "preserve"}, + args: []string{"openai"}, + expected: map[string]any{ + "existing-key": "preserve", + contractsai.OptionProvider: "openai", + }, + }, + { + name: "overrides previous value", + initial: map[string]any{ + contractsai.OptionProvider: "initial-provider", + }, + args: []string{"openai", "anthropic"}, + expected: map[string]any{ + contractsai.OptionProvider: "anthropic", + }, + }, + { + name: "panics on nil map", + args: []string{"openai"}, + nilMap: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if tt.nilMap { + assert.PanicsWithError(t, "assignment to entry in nil map", func() { + for _, arg := range tt.args { + WithProvider(arg)(nil) + } + }) + return + } + for _, arg := range tt.args { + WithProvider(arg)(tt.initial) + } + assert.Equal(t, tt.expected, tt.initial) + }) + } +} + +func TestWithModel(t *testing.T) { + tests := []struct { + name string + initial map[string]any + args []string + expected map[string]any + nilMap bool + }{ + { + name: "sets model while preserving existing keys", + initial: map[string]any{"existing-key": "preserve"}, + args: []string{"gpt-4"}, + expected: map[string]any{ + "existing-key": "preserve", + contractsai.OptionModel: "gpt-4", + }, + }, + { + name: "overrides previous value", + initial: map[string]any{ + contractsai.OptionModel: "initial-model", + }, + args: []string{"gpt-4", "gpt-4o"}, + expected: map[string]any{ + contractsai.OptionModel: "gpt-4o", + }, + }, + { + name: "panics on nil map", + args: []string{"gpt-4"}, + nilMap: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if tt.nilMap { + assert.PanicsWithError(t, "assignment to entry in nil map", func() { + for _, arg := range tt.args { + WithModel(arg)(nil) + } + }) + return + } + for _, arg := range tt.args { + WithModel(arg)(tt.initial) + } + assert.Equal(t, tt.expected, tt.initial) + }) + } +} diff --git a/ai/provider.go b/ai/provider.go new file mode 100644 index 000000000..ba3828fc2 --- /dev/null +++ b/ai/provider.go @@ -0,0 +1,63 @@ +package ai + +import ( + "sync" + + contractsai "github.com/goravel/framework/contracts/ai" + "github.com/goravel/framework/errors" +) + +type ProviderResolver struct { + config contractsai.Config + providers map[string]contractsai.Provider + mu sync.RWMutex +} + +func NewProviderResolver(config contractsai.Config) *ProviderResolver { + return &ProviderResolver{ + config: config, + providers: make(map[string]contractsai.Provider), + } +} + +func (r *ProviderResolver) New(providerName string) (contractsai.Provider, error) { + r.mu.RLock() + if provider, ok := r.providers[providerName]; ok { + r.mu.RUnlock() + return provider, nil + } + r.mu.RUnlock() + + providerCfg, ok := r.config.Providers[providerName] + if !ok { + return nil, errors.AIProviderNotSupported.Args(providerName) + } + + r.mu.Lock() + defer r.mu.Unlock() + + // Double-check after acquiring the write lock to avoid TOCTOU races. + if provider, ok := r.providers[providerName]; ok { + return provider, nil + } + + provider, err := r.resolve(providerName, providerCfg) + if err != nil { + return nil, err + } + if provider != nil { + r.providers[providerName] = provider + } + + return provider, nil +} + +func (r *ProviderResolver) resolve(name string, config contractsai.ProviderConfig) (contractsai.Provider, error) { + if p, ok := config.Via.(contractsai.Provider); ok { + return p, nil + } + if fn, ok := config.Via.(func() (contractsai.Provider, error)); ok { + return fn() + } + return nil, errors.AIProviderContractNotFulfilled.Args(name) +} diff --git a/ai/provider_test.go b/ai/provider_test.go new file mode 100644 index 000000000..e5f734f90 --- /dev/null +++ b/ai/provider_test.go @@ -0,0 +1,167 @@ +package ai + +import ( + "context" + "sync" + "testing" + + "github.com/stretchr/testify/assert" + + contractsai "github.com/goravel/framework/contracts/ai" + "github.com/goravel/framework/errors" +) + +type testProvider struct { + id string +} + +func (t *testProvider) Prompt(ctx context.Context, prompt contractsai.AgentPrompt) (contractsai.Response, error) { + return nil, nil +} + +func TestProviderResolver_New(t *testing.T) { + direct := &testProvider{id: "direct"} + + tests := []struct { + name string + config contractsai.Config + providerName string + wantProvider contractsai.Provider + wantErr error + }{ + { + name: "unsupported provider", + config: contractsai.Config{Providers: map[string]contractsai.ProviderConfig{}}, + providerName: "missing", + wantErr: errors.AIProviderNotSupported.Args("missing"), + }, + { + name: "direct provider via instance", + config: contractsai.Config{Providers: map[string]contractsai.ProviderConfig{ + "direct": {Via: direct}, + }}, + providerName: "direct", + wantProvider: direct, + }, + { + name: "contract not fulfilled", + config: contractsai.Config{Providers: map[string]contractsai.ProviderConfig{ + "broken": {Via: "not-a-provider"}, + }}, + providerName: "broken", + wantErr: errors.AIProviderContractNotFulfilled.Args("broken"), + }, + { + name: "factory returns error", + config: contractsai.Config{Providers: map[string]contractsai.ProviderConfig{ + "factory": {Via: func() (contractsai.Provider, error) { + return nil, assert.AnError + }}, + }}, + providerName: "factory", + wantErr: assert.AnError, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + resolver := NewProviderResolver(tt.config) + + got, err := resolver.New(tt.providerName) + + assert.Equal(t, tt.wantErr, err) + assert.Equal(t, tt.wantProvider, got) + }) + } +} + +func TestProviderResolver_NewCacheBehavior(t *testing.T) { + tests := []struct { + name string + setup func() (*ProviderResolver, contractsai.Provider, func() int) + wantFactoryCalled int + wantErr error + }{ + { + name: "successful factory provider is cached", + setup: func() (*ProviderResolver, contractsai.Provider, func() int) { + factoryCalled := 0 + want := &testProvider{id: "factory"} + resolver := NewProviderResolver(contractsai.Config{Providers: map[string]contractsai.ProviderConfig{ + "factory": {Via: func() (contractsai.Provider, error) { + factoryCalled++ + return want, nil + }}, + }}) + return resolver, want, func() int { return factoryCalled } + }, + wantFactoryCalled: 1, + }, + { + name: "failed factory provider is not cached", + setup: func() (*ProviderResolver, contractsai.Provider, func() int) { + factoryCalled := 0 + resolver := NewProviderResolver(contractsai.Config{Providers: map[string]contractsai.ProviderConfig{ + "factory": {Via: func() (contractsai.Provider, error) { + factoryCalled++ + return nil, assert.AnError + }}, + }}) + return resolver, nil, func() int { return factoryCalled } + }, + wantFactoryCalled: 2, + wantErr: assert.AnError, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + resolver, want, getFactoryCalled := tt.setup() + + first, firstErr := resolver.New("factory") + second, secondErr := resolver.New("factory") + + assert.Equal(t, tt.wantErr, firstErr) + assert.Equal(t, tt.wantErr, secondErr) + assert.Equal(t, want, first) + assert.Equal(t, want, second) + assert.Equal(t, tt.wantFactoryCalled, getFactoryCalled()) + }) + } +} + +func TestProviderResolver_NewConcurrent(t *testing.T) { + const goroutines = 50 + + var factoryCalled int + var mu sync.Mutex + want := &testProvider{id: "concurrent"} + + resolver := NewProviderResolver(contractsai.Config{Providers: map[string]contractsai.ProviderConfig{ + "p": {Via: func() (contractsai.Provider, error) { + mu.Lock() + factoryCalled++ + mu.Unlock() + return want, nil + }}, + }}) + + results := make([]contractsai.Provider, goroutines) + errs := make([]error, goroutines) + + var wg sync.WaitGroup + wg.Add(goroutines) + for i := range goroutines { + go func(i int) { + defer wg.Done() + results[i], errs[i] = resolver.New("p") + }(i) + } + wg.Wait() + + assert.Equal(t, 1, factoryCalled, "factory should be called exactly once") + for i := range goroutines { + assert.NoError(t, errs[i]) + assert.Equal(t, want, results[i]) + } +} diff --git a/ai/service_provider.go b/ai/service_provider.go index e03e40432..c5475cbed 100644 --- a/ai/service_provider.go +++ b/ai/service_provider.go @@ -1,6 +1,9 @@ package ai import ( + "context" + + contractsai "github.com/goravel/framework/contracts/ai" "github.com/goravel/framework/contracts/binding" "github.com/goravel/framework/contracts/foundation" ) @@ -18,7 +21,13 @@ func (r *ServiceProvider) Relationship() binding.Relationship { func (r *ServiceProvider) Register(app foundation.Application) { app.Singleton(binding.AI, func(app foundation.Application) (any, error) { - return NewApplication(app.MakeConfig()), nil + config := app.MakeConfig() + var aiConfig contractsai.Config + if err := config.UnmarshalKey("ai", &aiConfig); err != nil { + return nil, err + } + + return NewApplication(context.Background(), aiConfig), nil }) } diff --git a/ai/service_provider_test.go b/ai/service_provider_test.go new file mode 100644 index 000000000..15a96b92a --- /dev/null +++ b/ai/service_provider_test.go @@ -0,0 +1,113 @@ +package ai + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" + "github.com/stretchr/testify/suite" + + contractsai "github.com/goravel/framework/contracts/ai" + "github.com/goravel/framework/contracts/binding" + contractsfoundation "github.com/goravel/framework/contracts/foundation" + mocksconfig "github.com/goravel/framework/mocks/config" + mocksfoundation "github.com/goravel/framework/mocks/foundation" +) + +type ServiceProviderTestSuite struct { + suite.Suite +} + +func TestServiceProviderTestSuite(t *testing.T) { + suite.Run(t, &ServiceProviderTestSuite{}) +} + +func (s *ServiceProviderTestSuite) TestRelationship() { + provider := &ServiceProvider{} + + relationship := provider.Relationship() + s.Equal(binding.Relationship{ + Bindings: []string{binding.AI}, + Dependencies: binding.Bindings[binding.AI].Dependencies, + }, relationship) +} + +func (s *ServiceProviderTestSuite) TestRegister() { + var ( + mockApp *mocksfoundation.Application + mockCallbackApp *mocksfoundation.Application + mockConfig *mocksconfig.Config + callback func(contractsfoundation.Application) (any, error) + ) + + beforeEach := func() { + mockApp = mocksfoundation.NewApplication(s.T()) + mockCallbackApp = mocksfoundation.NewApplication(s.T()) + mockConfig = mocksconfig.NewConfig(s.T()) + + provider := &ServiceProvider{} + mockApp.EXPECT().Singleton(binding.AI, mock.MatchedBy(func(cb any) bool { + typedCallback, ok := cb.(func(contractsfoundation.Application) (any, error)) + if !ok { + return false + } + callback = typedCallback + return true + })).Once() + provider.Register(mockApp) + s.Require().NotNil(callback) + } + + tests := []struct { + name string + setup func() + expectError error + }{ + { + name: "binds application", + setup: func() { + mockCallbackApp.EXPECT().MakeConfig().Return(mockConfig).Once() + mockConfig.EXPECT(). + UnmarshalKey("ai", mock.MatchedBy(func(rawVal any) bool { + _, ok := rawVal.(*contractsai.Config) + return ok + })). + RunAndReturn(func(_ string, rawVal any) error { + *rawVal.(*contractsai.Config) = contractsai.Config{Default: "default"} + return nil + }). + Once() + }, + }, + { + name: "returns error when config cannot be unmarshaled", + setup: func() { + mockCallbackApp.EXPECT().MakeConfig().Return(mockConfig).Once() + mockConfig.EXPECT(). + UnmarshalKey("ai", mock.MatchedBy(func(rawVal any) bool { + _, ok := rawVal.(*contractsai.Config) + return ok + })). + Return(assert.AnError). + Once() + }, + expectError: assert.AnError, + }, + } + + for _, tt := range tests { + s.Run(tt.name, func() { + beforeEach() + tt.setup() + + instance, err := callback(mockCallbackApp) + s.Equal(tt.expectError, err) + if tt.expectError != nil { + s.Nil(instance) + return + } + s.IsType(&Application{}, instance) + s.Equal(contractsai.Config{Default: "default"}, instance.(*Application).config) + }) + } +} diff --git a/contracts/ai/ai.go b/contracts/ai/ai.go index c3696b801..f38789b7e 100644 --- a/contracts/ai/ai.go +++ b/contracts/ai/ai.go @@ -6,12 +6,14 @@ import "context" type AI interface { // Agent creates a conversation bound to the resolved driver. Agent(agent Agent, options ...Option) (Conversation, error) + // WithContext returns a new AI instance that carries the provided context for all operations. + WithContext(ctx context.Context) AI } // Conversation is a stateful chat session. type Conversation interface { // Prompt sends a non-streaming input and updates the conversation history. - Prompt(ctx context.Context, input string) (Response, error) + Prompt(input string) (Response, error) // Messages returns current conversation history. Messages() []Message // Reset clears runtime history and restores initial agent messages. diff --git a/contracts/ai/config.go b/contracts/ai/config.go new file mode 100644 index 000000000..d6730ac39 --- /dev/null +++ b/contracts/ai/config.go @@ -0,0 +1,19 @@ +package ai + +type Config struct { + Default string + Providers map[string]ProviderConfig +} + +type ProviderConfig struct { + Key string + Models ModelsConfig + Url string + Via any // Provider or func() (Provider, error) +} + +type ModelsConfig struct { + Text struct { + Default string + } +} diff --git a/contracts/ai/provider.go b/contracts/ai/provider.go index 8a143334c..de4cc7689 100644 --- a/contracts/ai/provider.go +++ b/contracts/ai/provider.go @@ -2,9 +2,17 @@ package ai import "context" -// Provider defines low-level model interactions. +// AgentPrompt carries all inputs the provider needs to call the model. +// Agent.Instructions() returns the system prompt; Agent.Messages() returns the runtime conversation history. +type AgentPrompt struct { + Agent Agent + Input string + Model string +} + +// Provider defines low-level model interactions (text generation). +// Future: extend with TextProvider, ImageProvider, AudioProvider, etc. type Provider interface { // Prompt executes a non-streaming model request. - // TODO: Optimize the parameters when implementing a real provider. - Prompt(ctx context.Context) (Response, error) + Prompt(ctx context.Context, prompt AgentPrompt) (Response, error) } diff --git a/errors/list.go b/errors/list.go index 49208ef8c..db74c9372 100644 --- a/errors/list.go +++ b/errors/list.go @@ -33,6 +33,9 @@ var ( AuthProviderDriverNotFound = New("driver %s for user provider %s was not found") AuthUnsupportedDriverMethod = New("The method was not supported for the driver %s") + AIProviderNotSupported = New("ai provider not found: %s") + AIProviderContractNotFulfilled = New("%s.via must be contracts/ai.Provider or func() (contracts/ai.Provider, error)") + CacheDriverNotSupported = New("invalid driver: %s, only support memory, custom") CacheForeverFailed = New("cache forever is failed") CacheMemoryDriverNotSupportDocker = New("memory driver doesn't support docker") diff --git a/go.mod b/go.mod index 2f1157071..0abbfc133 100644 --- a/go.mod +++ b/go.mod @@ -74,9 +74,14 @@ require ( github.com/mattn/go-colorable v0.1.14 // indirect github.com/mitchellh/go-homedir v1.1.0 // indirect github.com/mitchellh/mapstructure v1.5.0 // indirect + github.com/openai/openai-go v1.12.0 // indirect github.com/rs/zerolog v1.33.0 // indirect github.com/samber/slog-common v0.20.0 // indirect github.com/spf13/cobra v1.8.1 // indirect + github.com/tidwall/gjson v1.14.4 // indirect + github.com/tidwall/match v1.1.1 // indirect + github.com/tidwall/pretty v1.2.1 // indirect + github.com/tidwall/sjson v1.2.5 // indirect github.com/vektra/mockery/v2 v2.53.5 // indirect go.opentelemetry.io/auto/sdk v1.2.1 // indirect go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.42.0 // indirect diff --git a/go.sum b/go.sum index e32ea33ec..f9c9cdf37 100644 --- a/go.sum +++ b/go.sum @@ -194,6 +194,8 @@ github.com/muesli/cancelreader v0.2.2 h1:3I4Kt4BQjOR54NavqnDogx/MIoWBFa0StPA8ELU github.com/muesli/cancelreader v0.2.2/go.mod h1:3XuTXfFS2VjM+HTLZY9Ak0l6eUKfijIfMUZ4EgX0QYo= github.com/muesli/termenv v0.16.0 h1:S5AlUN9dENB57rsbnkPyfdGuWIlkmzJjbFf0Tf5FWUc= github.com/muesli/termenv v0.16.0/go.mod h1:ZRfOIKPFDYQoDFF4Olj7/QJbW60Ol/kL1pU3VfY/Cnk= +github.com/openai/openai-go v1.12.0 h1:NBQCnXzqOTv5wsgNC36PrFEiskGfO5wccfCWDo9S1U0= +github.com/openai/openai-go v1.12.0/go.mod h1:g461MYGXEXBVdV5SaR/5tNzNbSfwTBBefwc+LlDCK0Y= github.com/pelletier/go-toml/v2 v2.2.4 h1:mye9XuhQ6gvn5h28+VilKrrPoQVanw5PMw/TB0t5Ec4= github.com/pelletier/go-toml/v2 v2.2.4/go.mod h1:2gIqNv+qfxSVS7cM2xJQKtLSTLUE9V8t9Stt+h56mCY= github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0= @@ -256,6 +258,16 @@ github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= github.com/subosito/gotenv v1.6.0 h1:9NlTDc1FTs4qu0DDq7AEtTPNw6SVm7uBMsUCUjABIf8= github.com/subosito/gotenv v1.6.0/go.mod h1:Dk4QP5c2W3ibzajGcXpNraDfq2IrhjMIvMSWPKKo0FU= +github.com/tidwall/gjson v1.14.2/go.mod h1:/wbyibRr2FHMks5tjHJ5F8dMZh3AcwJEMf5vlfC0lxk= +github.com/tidwall/gjson v1.14.4 h1:uo0p8EbA09J7RQaflQ1aBRffTR7xedD2bcIVSYxLnkM= +github.com/tidwall/gjson v1.14.4/go.mod h1:/wbyibRr2FHMks5tjHJ5F8dMZh3AcwJEMf5vlfC0lxk= +github.com/tidwall/match v1.1.1 h1:+Ho715JplO36QYgwN9PGYNhgZvoUSc9X2c80KVTi+GA= +github.com/tidwall/match v1.1.1/go.mod h1:eRSPERbgtNPcGhD8UCthc6PmLEQXEWd3PRB5JTxsfmM= +github.com/tidwall/pretty v1.2.0/go.mod h1:ITEVvHYasfjBbM0u2Pg8T2nJnzm8xPwvNhhsoaGGjNU= +github.com/tidwall/pretty v1.2.1 h1:qjsOFOWWQl+N3RsoF5/ssm1pHmJJwhjlSbZ51I6wMl4= +github.com/tidwall/pretty v1.2.1/go.mod h1:ITEVvHYasfjBbM0u2Pg8T2nJnzm8xPwvNhhsoaGGjNU= +github.com/tidwall/sjson v1.2.5 h1:kLy8mja+1c9jlljvWTlSazM7cKDRfJuR/bOJhcY5NcY= +github.com/tidwall/sjson v1.2.5/go.mod h1:Fvgq9kS/6ociJEDnK0Fk1cpYF4FIW6ZF7LAe+6jwd28= github.com/urfave/cli/v3 v3.7.0 h1:AGSnbUyjtLiM+WJUb4dzXKldl/gL+F8OwmRDtVr6g2U= github.com/urfave/cli/v3 v3.7.0/go.mod h1:ysVLtOEmg2tOy6PknnYVhDoouyC/6N42TMeoMzskhso= github.com/vektra/mockery/v2 v2.53.5 h1:iktAY68pNiMvLoHxKqlSNSv/1py0QF/17UGrrAMYDI8= diff --git a/mocks/ai/AI.go b/mocks/ai/AI.go index 1dd529fcd..9daf4a6bf 100644 --- a/mocks/ai/AI.go +++ b/mocks/ai/AI.go @@ -3,7 +3,10 @@ package ai import ( + context "context" + ai "github.com/goravel/framework/contracts/ai" + mock "github.com/stretchr/testify/mock" ) @@ -93,6 +96,54 @@ func (_c *AI_Agent_Call) RunAndReturn(run func(ai.Agent, ...ai.Option) (ai.Conve return _c } +// WithContext provides a mock function with given fields: ctx +func (_m *AI) WithContext(ctx context.Context) ai.AI { + ret := _m.Called(ctx) + + if len(ret) == 0 { + panic("no return value specified for WithContext") + } + + var r0 ai.AI + if rf, ok := ret.Get(0).(func(context.Context) ai.AI); ok { + r0 = rf(ctx) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(ai.AI) + } + } + + return r0 +} + +// AI_WithContext_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'WithContext' +type AI_WithContext_Call struct { + *mock.Call +} + +// WithContext is a helper method to define mock.On call +// - ctx context.Context +func (_e *AI_Expecter) WithContext(ctx interface{}) *AI_WithContext_Call { + return &AI_WithContext_Call{Call: _e.mock.On("WithContext", ctx)} +} + +func (_c *AI_WithContext_Call) Run(run func(ctx context.Context)) *AI_WithContext_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(context.Context)) + }) + return _c +} + +func (_c *AI_WithContext_Call) Return(_a0 ai.AI) *AI_WithContext_Call { + _c.Call.Return(_a0) + return _c +} + +func (_c *AI_WithContext_Call) RunAndReturn(run func(context.Context) ai.AI) *AI_WithContext_Call { + _c.Call.Return(run) + return _c +} + // NewAI creates a new instance of AI. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. // The first argument is typically a *testing.T value. func NewAI(t interface { diff --git a/mocks/ai/Conversation.go b/mocks/ai/Conversation.go index 65ba1a744..820f2ed88 100644 --- a/mocks/ai/Conversation.go +++ b/mocks/ai/Conversation.go @@ -3,10 +3,7 @@ package ai import ( - context "context" - ai "github.com/goravel/framework/contracts/ai" - mock "github.com/stretchr/testify/mock" ) @@ -70,9 +67,9 @@ func (_c *Conversation_Messages_Call) RunAndReturn(run func() []ai.Message) *Con return _c } -// Prompt provides a mock function with given fields: ctx, input -func (_m *Conversation) Prompt(ctx context.Context, input string) (ai.Response, error) { - ret := _m.Called(ctx, input) +// Prompt provides a mock function with given fields: input +func (_m *Conversation) Prompt(input string) (ai.Response, error) { + ret := _m.Called(input) if len(ret) == 0 { panic("no return value specified for Prompt") @@ -80,19 +77,19 @@ func (_m *Conversation) Prompt(ctx context.Context, input string) (ai.Response, var r0 ai.Response var r1 error - if rf, ok := ret.Get(0).(func(context.Context, string) (ai.Response, error)); ok { - return rf(ctx, input) + if rf, ok := ret.Get(0).(func(string) (ai.Response, error)); ok { + return rf(input) } - if rf, ok := ret.Get(0).(func(context.Context, string) ai.Response); ok { - r0 = rf(ctx, input) + if rf, ok := ret.Get(0).(func(string) ai.Response); ok { + r0 = rf(input) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(ai.Response) } } - if rf, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = rf(ctx, input) + if rf, ok := ret.Get(1).(func(string) error); ok { + r1 = rf(input) } else { r1 = ret.Error(1) } @@ -106,15 +103,14 @@ type Conversation_Prompt_Call struct { } // Prompt is a helper method to define mock.On call -// - ctx context.Context // - input string -func (_e *Conversation_Expecter) Prompt(ctx interface{}, input interface{}) *Conversation_Prompt_Call { - return &Conversation_Prompt_Call{Call: _e.mock.On("Prompt", ctx, input)} +func (_e *Conversation_Expecter) Prompt(input interface{}) *Conversation_Prompt_Call { + return &Conversation_Prompt_Call{Call: _e.mock.On("Prompt", input)} } -func (_c *Conversation_Prompt_Call) Run(run func(ctx context.Context, input string)) *Conversation_Prompt_Call { +func (_c *Conversation_Prompt_Call) Run(run func(input string)) *Conversation_Prompt_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].(context.Context), args[1].(string)) + run(args[0].(string)) }) return _c } @@ -124,7 +120,7 @@ func (_c *Conversation_Prompt_Call) Return(_a0 ai.Response, _a1 error) *Conversa return _c } -func (_c *Conversation_Prompt_Call) RunAndReturn(run func(context.Context, string) (ai.Response, error)) *Conversation_Prompt_Call { +func (_c *Conversation_Prompt_Call) RunAndReturn(run func(string) (ai.Response, error)) *Conversation_Prompt_Call { _c.Call.Return(run) return _c } diff --git a/mocks/ai/Provider.go b/mocks/ai/Provider.go index 322fc04fd..31c6f5ee1 100644 --- a/mocks/ai/Provider.go +++ b/mocks/ai/Provider.go @@ -23,9 +23,9 @@ func (_m *Provider) EXPECT() *Provider_Expecter { return &Provider_Expecter{mock: &_m.Mock} } -// Prompt provides a mock function with given fields: ctx -func (_m *Provider) Prompt(ctx context.Context) (ai.Response, error) { - ret := _m.Called(ctx) +// Prompt provides a mock function with given fields: ctx, prompt +func (_m *Provider) Prompt(ctx context.Context, prompt ai.AgentPrompt) (ai.Response, error) { + ret := _m.Called(ctx, prompt) if len(ret) == 0 { panic("no return value specified for Prompt") @@ -33,19 +33,19 @@ func (_m *Provider) Prompt(ctx context.Context) (ai.Response, error) { var r0 ai.Response var r1 error - if rf, ok := ret.Get(0).(func(context.Context) (ai.Response, error)); ok { - return rf(ctx) + if rf, ok := ret.Get(0).(func(context.Context, ai.AgentPrompt) (ai.Response, error)); ok { + return rf(ctx, prompt) } - if rf, ok := ret.Get(0).(func(context.Context) ai.Response); ok { - r0 = rf(ctx) + if rf, ok := ret.Get(0).(func(context.Context, ai.AgentPrompt) ai.Response); ok { + r0 = rf(ctx, prompt) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(ai.Response) } } - if rf, ok := ret.Get(1).(func(context.Context) error); ok { - r1 = rf(ctx) + if rf, ok := ret.Get(1).(func(context.Context, ai.AgentPrompt) error); ok { + r1 = rf(ctx, prompt) } else { r1 = ret.Error(1) } @@ -60,13 +60,14 @@ type Provider_Prompt_Call struct { // Prompt is a helper method to define mock.On call // - ctx context.Context -func (_e *Provider_Expecter) Prompt(ctx interface{}) *Provider_Prompt_Call { - return &Provider_Prompt_Call{Call: _e.mock.On("Prompt", ctx)} +// - prompt ai.AgentPrompt +func (_e *Provider_Expecter) Prompt(ctx interface{}, prompt interface{}) *Provider_Prompt_Call { + return &Provider_Prompt_Call{Call: _e.mock.On("Prompt", ctx, prompt)} } -func (_c *Provider_Prompt_Call) Run(run func(ctx context.Context)) *Provider_Prompt_Call { +func (_c *Provider_Prompt_Call) Run(run func(ctx context.Context, prompt ai.AgentPrompt)) *Provider_Prompt_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].(context.Context)) + run(args[0].(context.Context), args[1].(ai.AgentPrompt)) }) return _c } @@ -76,7 +77,7 @@ func (_c *Provider_Prompt_Call) Return(_a0 ai.Response, _a1 error) *Provider_Pro return _c } -func (_c *Provider_Prompt_Call) RunAndReturn(run func(context.Context) (ai.Response, error)) *Provider_Prompt_Call { +func (_c *Provider_Prompt_Call) RunAndReturn(run func(context.Context, ai.AgentPrompt) (ai.Response, error)) *Provider_Prompt_Call { _c.Call.Return(run) return _c } diff --git a/mocks/ai/ProviderCreator.go b/mocks/ai/ProviderCreator.go deleted file mode 100644 index 1e2b0cc17..000000000 --- a/mocks/ai/ProviderCreator.go +++ /dev/null @@ -1,96 +0,0 @@ -// Code generated by mockery. DO NOT EDIT. - -package ai - -import ( - context "context" - - ai "github.com/goravel/framework/contracts/ai" - - mock "github.com/stretchr/testify/mock" -) - -// ProviderCreator is an autogenerated mock type for the ProviderCreator type -type ProviderCreator struct { - mock.Mock -} - -type ProviderCreator_Expecter struct { - mock *mock.Mock -} - -func (_m *ProviderCreator) EXPECT() *ProviderCreator_Expecter { - return &ProviderCreator_Expecter{mock: &_m.Mock} -} - -// Execute provides a mock function with given fields: ctx -func (_m *ProviderCreator) Execute(ctx context.Context) (ai.Provider, error) { - ret := _m.Called(ctx) - - if len(ret) == 0 { - panic("no return value specified for Execute") - } - - var r0 ai.Provider - var r1 error - if rf, ok := ret.Get(0).(func(context.Context) (ai.Provider, error)); ok { - return rf(ctx) - } - if rf, ok := ret.Get(0).(func(context.Context) ai.Provider); ok { - r0 = rf(ctx) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(ai.Provider) - } - } - - if rf, ok := ret.Get(1).(func(context.Context) error); ok { - r1 = rf(ctx) - } else { - r1 = ret.Error(1) - } - - return r0, r1 -} - -// ProviderCreator_Execute_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Execute' -type ProviderCreator_Execute_Call struct { - *mock.Call -} - -// Execute is a helper method to define mock.On call -// - ctx context.Context -func (_e *ProviderCreator_Expecter) Execute(ctx interface{}) *ProviderCreator_Execute_Call { - return &ProviderCreator_Execute_Call{Call: _e.mock.On("Execute", ctx)} -} - -func (_c *ProviderCreator_Execute_Call) Run(run func(ctx context.Context)) *ProviderCreator_Execute_Call { - _c.Call.Run(func(args mock.Arguments) { - run(args[0].(context.Context)) - }) - return _c -} - -func (_c *ProviderCreator_Execute_Call) Return(_a0 ai.Provider, _a1 error) *ProviderCreator_Execute_Call { - _c.Call.Return(_a0, _a1) - return _c -} - -func (_c *ProviderCreator_Execute_Call) RunAndReturn(run func(context.Context) (ai.Provider, error)) *ProviderCreator_Execute_Call { - _c.Call.Return(run) - return _c -} - -// NewProviderCreator creates a new instance of ProviderCreator. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewProviderCreator(t interface { - mock.TestingT - Cleanup(func()) -}) *ProviderCreator { - mock := &ProviderCreator{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -}