Compare commits
2 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 47671f3098 | |||
| 28195519d1 |
@@ -74,9 +74,10 @@ responses OpenAI Responses API /v1/responses
|
|||||||
```
|
```
|
||||||
|
|
||||||
The Responses API client converts the conversation history to Responses input
|
The Responses API client converts the conversation history to Responses input
|
||||||
items (`user`/`assistant`/`system` messages, `function_call`, and
|
items (`user`/`assistant` messages, `function_call`, and
|
||||||
`function_call_output`), streams `response.output_text.delta` events, and
|
`function_call_output`), sends the system prompt as the top-level
|
||||||
supports function calling through `response.output_item.added`,
|
`instructions` field, streams `response.output_text.delta` events, and supports
|
||||||
|
function calling through `response.output_item.added`,
|
||||||
`response.function_call_arguments.delta`, and `response.output_item.done`.
|
`response.function_call_arguments.delta`, and `response.output_item.done`.
|
||||||
|
|
||||||
`models` is required. `/model model <name>` is restricted to that list. Each
|
`models` is required. `/model model <name>` is restricted to that list. Each
|
||||||
@@ -98,10 +99,11 @@ none, middle, high, xhigh, max
|
|||||||
```
|
```
|
||||||
|
|
||||||
When configured or changed with `/model thinking ...`, the value is sent as a
|
When configured or changed with `/model thinking ...`, the value is sent as a
|
||||||
top-level request field named `thinking`. For OpenAI reasoning models on the
|
top-level request field. The default field depends on `api_type`: `chat` uses
|
||||||
Responses API, set `thinking_param: reasoning` so the level is sent as the
|
`thinking`, while `responses` uses the official `reasoning` field, sent as
|
||||||
official `reasoning: {"effort": ...}` field; agentu maps its `middle` level to
|
`{"effort": ...}` with agentu's `middle` level mapped to the official `medium`
|
||||||
the official `medium` value and passes `none`, `high`, `xhigh`, `max` through.
|
value (`none`, `high`, `xhigh`, `max` pass through). Override with
|
||||||
|
`thinking_param` if a provider expects a different field name.
|
||||||
|
|
||||||
## CLI Flags
|
## CLI Flags
|
||||||
|
|
||||||
@@ -165,7 +167,8 @@ Web tools are also available in read-only mode:
|
|||||||
|
|
||||||
- `file_write` for creating or replacing whole files
|
- `file_write` for creating or replacing whole files
|
||||||
- `file_edit` for surgical line-range or pattern edits to existing files
|
- `file_edit` for surgical line-range or pattern edits to existing files
|
||||||
- `shell_run` for non-interactive shell commands
|
- `shell_run` for non-interactive shell commands; output streams live into the
|
||||||
|
TUI while the command runs instead of appearing only after completion
|
||||||
|
|
||||||
Use `--yolo` for prompts that ask agentu to run or inspect local commands, for
|
Use `--yolo` for prompts that ask agentu to run or inspect local commands, for
|
||||||
example:
|
example:
|
||||||
|
|||||||
+56
-12
@@ -24,6 +24,7 @@ const (
|
|||||||
toolLogArgumentLimit = 400
|
toolLogArgumentLimit = 400
|
||||||
toolLogOutputPreviewBytes = 4096
|
toolLogOutputPreviewBytes = 4096
|
||||||
toolLogOutputMarker = "[output]"
|
toolLogOutputMarker = "[output]"
|
||||||
|
toolStreamLogLimit = tools.MaxToolOutputBytes
|
||||||
|
|
||||||
summarizerSystemPrompt = "You are a conversation summarizer. Summarize the following conversation history into a concise but detailed summary. Preserve:\n" +
|
summarizerSystemPrompt = "You are a conversation summarizer. Summarize the following conversation history into a concise but detailed summary. Preserve:\n" +
|
||||||
"- Key facts, decisions, and conclusions\n" +
|
"- Key facts, decisions, and conclusions\n" +
|
||||||
@@ -182,6 +183,7 @@ func (a *Agent) SetLastUsage(usage *llm.Usage) {
|
|||||||
copy := *usage
|
copy := *usage
|
||||||
a.lastUsage = ©
|
a.lastUsage = ©
|
||||||
}
|
}
|
||||||
|
|
||||||
// compactRetainedTokenBudget returns the number of estimated tokens to retain
|
// compactRetainedTokenBudget returns the number of estimated tokens to retain
|
||||||
// based on compactRecentRetentionPercent, using integer ceiling division.
|
// based on compactRecentRetentionPercent, using integer ceiling division.
|
||||||
func compactRetainedTokenBudget(totalTokens int) int {
|
func compactRetainedTokenBudget(totalTokens int) int {
|
||||||
@@ -237,7 +239,6 @@ func compactRecentStart(messages []llm.Message, prefixEnd int) int {
|
|||||||
return compactTurnBoundary(messages, prefixEnd, start)
|
return compactTurnBoundary(messages, prefixEnd, start)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
// Compact summarizes older conversation messages while preserving the system
|
// Compact summarizes older conversation messages while preserving the system
|
||||||
// prompt and recent messages intact.
|
// prompt and recent messages intact.
|
||||||
func (a *Agent) Compact(ctx context.Context) error {
|
func (a *Agent) Compact(ctx context.Context) error {
|
||||||
@@ -495,32 +496,75 @@ func (a *Agent) toolDefinitions() []llm.Tool {
|
|||||||
return a.toolRegistry.Definitions()
|
return a.toolRegistry.Definitions()
|
||||||
}
|
}
|
||||||
|
|
||||||
func (a *Agent) executeTool(ctx context.Context, call llm.ToolCall) string {
|
func (a *Agent) executeTool(ctx context.Context, call llm.ToolCall, logs io.Writer) (string, bool) {
|
||||||
name := call.Function.Name
|
name := call.Function.Name
|
||||||
|
|
||||||
if a.toolRegistry == nil {
|
if a.toolRegistry == nil {
|
||||||
return tools.FormatError(fmt.Errorf("tool registry is disabled"))
|
return tools.FormatError(fmt.Errorf("tool registry is disabled")), false
|
||||||
}
|
}
|
||||||
tool, ok := a.toolRegistry.Get(name)
|
tool, ok := a.toolRegistry.Get(name)
|
||||||
if !ok {
|
if !ok {
|
||||||
return tools.FormatError(fmt.Errorf("unknown tool: %s", name))
|
return tools.FormatError(fmt.Errorf("unknown tool: %s", name)), false
|
||||||
}
|
}
|
||||||
|
|
||||||
args := strings.TrimSpace(call.Function.Arguments)
|
args := strings.TrimSpace(call.Function.Arguments)
|
||||||
if args == "" {
|
if args == "" {
|
||||||
args = "{}"
|
args = "{}"
|
||||||
}
|
}
|
||||||
|
raw := json.RawMessage(args)
|
||||||
|
|
||||||
toolCtx, cancel := context.WithTimeout(ctx, a.toolTimeout)
|
toolCtx, cancel := context.WithTimeout(ctx, a.toolTimeout)
|
||||||
defer cancel()
|
defer cancel()
|
||||||
output, err := tool.Execute(toolCtx, json.RawMessage(args))
|
|
||||||
|
if streaming, ok := tool.(tools.StreamingTool); ok {
|
||||||
|
return a.executeStreamingTool(streaming, toolCtx, raw, name, args, logs)
|
||||||
|
}
|
||||||
|
|
||||||
|
output, err := tool.Execute(toolCtx, raw)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if output != "" {
|
if output != "" {
|
||||||
return output + "\n" + tools.FormatError(err)
|
return output + "\n" + tools.FormatError(err), false
|
||||||
}
|
}
|
||||||
return tools.FormatError(err)
|
return tools.FormatError(err), false
|
||||||
}
|
}
|
||||||
return output
|
return output, false
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *Agent) executeStreamingTool(streaming tools.StreamingTool, ctx context.Context, raw json.RawMessage, name, args string, logs io.Writer) (string, bool) {
|
||||||
|
var streamed int
|
||||||
|
emit := func(chunk string) error {
|
||||||
|
if chunk == "" || logs == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if streamed >= toolStreamLogLimit {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if len(chunk) > toolStreamLogLimit-streamed {
|
||||||
|
chunk = chunk[:toolStreamLogLimit-streamed]
|
||||||
|
streamed = toolStreamLogLimit
|
||||||
|
_, err := fmt.Fprintf(logs, "\n[tool] %s %s\n%s\n%s\n[streaming output truncated after %d bytes]", name, compact(args, toolLogArgumentLimit), toolLogOutputMarker, chunk, toolStreamLogLimit)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
streamed += len(chunk)
|
||||||
|
_, err := fmt.Fprintf(logs, "\n[tool] %s %s\n%s\n%s", name, compact(args, toolLogArgumentLimit), toolLogOutputMarker, chunk)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
output, err := streaming.ExecuteStream(ctx, raw, emit)
|
||||||
|
if logs != nil {
|
||||||
|
if err != nil {
|
||||||
|
_ = emit("\n" + tools.FormatError(err))
|
||||||
|
} else if streamed == 0 {
|
||||||
|
_ = emit("(no output)")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
if output != "" {
|
||||||
|
return output + "\n" + tools.FormatError(err), true
|
||||||
|
}
|
||||||
|
return tools.FormatError(err), true
|
||||||
|
}
|
||||||
|
return output, true
|
||||||
}
|
}
|
||||||
|
|
||||||
func (a *Agent) hasMutatingCall(calls []llm.ToolCall) bool {
|
func (a *Agent) hasMutatingCall(calls []llm.ToolCall) bool {
|
||||||
@@ -550,9 +594,9 @@ func (a *Agent) executeTools(ctx context.Context, calls []llm.ToolCall, logs io.
|
|||||||
if a.hasMutatingCall(calls) {
|
if a.hasMutatingCall(calls) {
|
||||||
// Sequential: avoid racing mutating tools.
|
// Sequential: avoid racing mutating tools.
|
||||||
for i, call := range calls {
|
for i, call := range calls {
|
||||||
output := a.executeTool(ctx, call)
|
output, streamed := a.executeTool(ctx, call, logs)
|
||||||
results[i] = toolExecutionResult{call: call, output: output}
|
results[i] = toolExecutionResult{call: call, output: output}
|
||||||
if logs != nil {
|
if logs != nil && !streamed {
|
||||||
logToolOutput(logs, call, output)
|
logToolOutput(logs, call, output)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -566,9 +610,9 @@ func (a *Agent) executeTools(ctx context.Context, calls []llm.ToolCall, logs io.
|
|||||||
wg.Add(1)
|
wg.Add(1)
|
||||||
go func(idx int, c llm.ToolCall) {
|
go func(idx int, c llm.ToolCall) {
|
||||||
defer wg.Done()
|
defer wg.Done()
|
||||||
output := a.executeTool(ctx, c)
|
output, streamed := a.executeTool(ctx, c, logs)
|
||||||
results[idx] = toolExecutionResult{call: c, output: output}
|
results[idx] = toolExecutionResult{call: c, output: output}
|
||||||
if logs != nil {
|
if logs != nil && !streamed {
|
||||||
mu.Lock()
|
mu.Lock()
|
||||||
logToolOutput(logs, c, output)
|
logToolOutput(logs, c, output)
|
||||||
mu.Unlock()
|
mu.Unlock()
|
||||||
|
|||||||
@@ -87,6 +87,86 @@ func TestAgentRunsToolLoop(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type streamToolProvider struct {
|
||||||
|
calls int
|
||||||
|
t *testing.T
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *streamToolProvider) ChatStream(ctx context.Context, req llm.ChatRequest, emit func(llm.StreamEvent) error) error {
|
||||||
|
p.calls++
|
||||||
|
switch p.calls {
|
||||||
|
case 1:
|
||||||
|
return emit(llm.StreamEvent{ToolCalls: []llm.ToolCallDelta{
|
||||||
|
{Index: 0, ID: "call_1", Type: "function", Name: "test_stream", Arguments: `{}`},
|
||||||
|
}})
|
||||||
|
case 2:
|
||||||
|
last := req.Messages[len(req.Messages)-1]
|
||||||
|
if last.Role != llm.RoleTool || last.Content != "full result" {
|
||||||
|
p.t.Fatalf("last message = %#v", last)
|
||||||
|
}
|
||||||
|
return emit(llm.StreamEvent{Content: "done"})
|
||||||
|
default:
|
||||||
|
p.t.Fatalf("unexpected call count %d", p.calls)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type streamEchoTool struct{}
|
||||||
|
|
||||||
|
func (streamEchoTool) Definition() llm.Tool {
|
||||||
|
return llm.Tool{
|
||||||
|
Type: "function",
|
||||||
|
Function: llm.ToolFunction{
|
||||||
|
Name: "test_stream",
|
||||||
|
Description: "streaming test",
|
||||||
|
Parameters: json.RawMessage(`{"type":"object"}`),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (streamEchoTool) Execute(ctx context.Context, raw json.RawMessage) (string, error) {
|
||||||
|
return "full result", nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (streamEchoTool) ExecuteStream(ctx context.Context, raw json.RawMessage, emit func(string) error) (string, error) {
|
||||||
|
if err := emit("one\n"); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
if err := emit("two\n"); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
return "full result", nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentStreamsToolOutputToLogs(t *testing.T) {
|
||||||
|
provider := &streamToolProvider{t: t}
|
||||||
|
a := New(Options{
|
||||||
|
Provider: provider,
|
||||||
|
Model: "test",
|
||||||
|
SystemPrompt: "system",
|
||||||
|
ToolRegistry: tools.NewRegistry(streamEchoTool{}),
|
||||||
|
})
|
||||||
|
|
||||||
|
var out strings.Builder
|
||||||
|
var logs strings.Builder
|
||||||
|
if err := a.RunTurn(context.Background(), "stream", &out, &logs); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if out.String() != "done" {
|
||||||
|
t.Fatalf("out = %q", out.String())
|
||||||
|
}
|
||||||
|
logText := logs.String()
|
||||||
|
if strings.Count(logText, "[output]") != 2 {
|
||||||
|
t.Fatalf("logs = %q", logText)
|
||||||
|
}
|
||||||
|
if !strings.Contains(logText, "one\n") || !strings.Contains(logText, "two\n") {
|
||||||
|
t.Fatalf("logs missing streamed chunks: %q", logText)
|
||||||
|
}
|
||||||
|
if strings.Contains(logText, "[output]\nfull result") {
|
||||||
|
t.Fatalf("logs contain duplicate full output: %q", logText)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
type contentProvider struct {
|
type contentProvider struct {
|
||||||
called bool
|
called bool
|
||||||
text string
|
text string
|
||||||
|
|||||||
@@ -68,11 +68,29 @@ func (m *Manager) Resume(id string) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
m.currentSession = sess
|
m.currentSession = sess
|
||||||
|
m.restoreRuntime(sess)
|
||||||
m.agent.SetMessages(sess.Messages)
|
m.agent.SetMessages(sess.Messages)
|
||||||
m.agent.SetLastUsage(sess.LastUsage)
|
m.agent.SetLastUsage(sess.LastUsage)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// restoreRuntime switches the manager and agent back to the provider, model,
|
||||||
|
// and thinking level recorded on a resumed session. Values that are no longer
|
||||||
|
// valid for the current config are skipped and the defaults remain active.
|
||||||
|
func (m *Manager) restoreRuntime(sess *Session) {
|
||||||
|
if sess.Provider != "" && sess.Provider != m.providerName {
|
||||||
|
if _, ok := m.cfg.Providers[sess.Provider]; ok {
|
||||||
|
_, _ = m.SetProvider(sess.Provider)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if sess.Model != "" && sess.Model != m.provider.Model {
|
||||||
|
_, _ = m.SetModel(sess.Model)
|
||||||
|
}
|
||||||
|
if sess.Thinking != "" && sess.Thinking != m.provider.Thinking {
|
||||||
|
_, _ = m.SetThinking(sess.Thinking)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// Save persists the current session to disk.
|
// Save persists the current session to disk.
|
||||||
func (m *Manager) Save() error {
|
func (m *Manager) Save() error {
|
||||||
if m.store == nil || m.currentSession == nil {
|
if m.store == nil || m.currentSession == nil {
|
||||||
@@ -82,6 +100,8 @@ func (m *Manager) Save() error {
|
|||||||
m.currentSession.LastUsage = m.agent.LastUsage()
|
m.currentSession.LastUsage = m.agent.LastUsage()
|
||||||
m.currentSession.Provider = m.providerName
|
m.currentSession.Provider = m.providerName
|
||||||
m.currentSession.Model = m.provider.Model
|
m.currentSession.Model = m.provider.Model
|
||||||
|
m.currentSession.APIType = m.provider.APIType
|
||||||
|
m.currentSession.Thinking = m.provider.Thinking
|
||||||
// Auto-name if empty
|
// Auto-name if empty
|
||||||
if m.currentSession.Name == "" {
|
if m.currentSession.Name == "" {
|
||||||
m.currentSession.Name = DefaultName(m.currentSession.Messages)
|
m.currentSession.Name = DefaultName(m.currentSession.Messages)
|
||||||
@@ -150,6 +170,8 @@ func (m *Manager) NewSession() (string, error) {
|
|||||||
m.currentSession.Messages = msgs
|
m.currentSession.Messages = msgs
|
||||||
m.currentSession.Provider = m.providerName
|
m.currentSession.Provider = m.providerName
|
||||||
m.currentSession.Model = m.provider.Model
|
m.currentSession.Model = m.provider.Model
|
||||||
|
m.currentSession.APIType = m.provider.APIType
|
||||||
|
m.currentSession.Thinking = m.provider.Thinking
|
||||||
if m.currentSession.Name == "" {
|
if m.currentSession.Name == "" {
|
||||||
m.currentSession.Name = DefaultName(msgs)
|
m.currentSession.Name = DefaultName(msgs)
|
||||||
}
|
}
|
||||||
@@ -257,7 +279,7 @@ func (m *Manager) SetThinking(level string) (string, error) {
|
|||||||
m.provider.Thinking = level
|
m.provider.Thinking = level
|
||||||
m.provider.ThinkingConfigured = true
|
m.provider.ThinkingConfigured = true
|
||||||
if m.provider.ThinkingParam == "" {
|
if m.provider.ThinkingParam == "" {
|
||||||
m.provider.ThinkingParam = "thinking"
|
m.provider.ThinkingParam = config.DefaultThinkingParam(m.provider.APIType)
|
||||||
}
|
}
|
||||||
m.cfg.Providers[m.providerName] = m.provider
|
m.cfg.Providers[m.providerName] = m.provider
|
||||||
m.syncAgent(nil)
|
m.syncAgent(nil)
|
||||||
|
|||||||
@@ -403,3 +403,132 @@ func TestManagerUsesModelContextForAgentCompaction(t *testing.T) {
|
|||||||
t.Fatalf("answerCalls = %d, want 1", provider.answerCalls)
|
t.Fatalf("answerCalls = %d, want 1", provider.answerCalls)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestManagerResumeRestoresRuntime(t *testing.T) {
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
w.Header().Set("Content-Type", "text/event-stream")
|
||||||
|
_, _ = w.Write([]byte(`data: {"type":"response.output_text.delta","item_id":"msg_1","output_index":0,"content_index":0,"delta":"ok"}` + "\n\n"))
|
||||||
|
_, _ = w.Write([]byte(`data: {"type":"response.completed","response":{"id":"resp_1","status":"completed"}}` + "\n\n"))
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
cfg := &config.Config{
|
||||||
|
Providers: map[string]config.ProviderConfig{
|
||||||
|
"chat": {
|
||||||
|
Name: "chat",
|
||||||
|
APIType: config.APITypeChat,
|
||||||
|
BaseURL: "https://chat.example.com",
|
||||||
|
APIKey: "sk-test",
|
||||||
|
Model: "model-a",
|
||||||
|
Models: []string{"model-a"},
|
||||||
|
Thinking: "none",
|
||||||
|
ThinkingParam: "thinking",
|
||||||
|
ThinkingConfigured: true,
|
||||||
|
},
|
||||||
|
"resp": {
|
||||||
|
Name: "resp",
|
||||||
|
APIType: config.APITypeResponses,
|
||||||
|
BaseURL: server.URL,
|
||||||
|
APIKey: "sk-test",
|
||||||
|
Model: "m1",
|
||||||
|
Models: []string{"m1", "m2"},
|
||||||
|
Thinking: "none",
|
||||||
|
ThinkingParam: "reasoning",
|
||||||
|
ThinkingConfigured: true,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
assistant := agent.New(agent.Options{
|
||||||
|
Provider: &recordingProvider{},
|
||||||
|
Model: "model-a",
|
||||||
|
Thinking: "none",
|
||||||
|
ThinkingParam: "thinking",
|
||||||
|
ThinkingEnabled: true,
|
||||||
|
})
|
||||||
|
store, err := NewStore(filepath.Join(t.TempDir(), "sessions"))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
manager, err := NewManager(cfg, assistant, server.Client(), store)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Build a session on the responses provider with a custom model and thinking.
|
||||||
|
manager.newSession()
|
||||||
|
if _, err := manager.SetProvider("resp"); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, err := manager.SetModel("m2"); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, err := manager.SetThinking("xhigh"); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := manager.Save(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
id := manager.CurrentSessionID()
|
||||||
|
|
||||||
|
// Drift away from the saved runtime.
|
||||||
|
if _, err := manager.SetProvider("chat"); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, err := manager.SetModel("model-a"); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, err := manager.SetThinking("none"); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := manager.Resume(id); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if got := manager.CurrentProvider(); got != "resp" {
|
||||||
|
t.Fatalf("provider = %q, want resp", got)
|
||||||
|
}
|
||||||
|
if got := manager.CurrentModel(); got != "m2" {
|
||||||
|
t.Fatalf("model = %q, want m2", got)
|
||||||
|
}
|
||||||
|
if got := manager.CurrentThinking(); got != "xhigh" {
|
||||||
|
t.Fatalf("thinking = %q, want xhigh", got)
|
||||||
|
}
|
||||||
|
if got := manager.CurrentAPIType(); got != config.APITypeResponses {
|
||||||
|
t.Fatalf("api type = %q, want responses", got)
|
||||||
|
}
|
||||||
|
runtime := assistant.RuntimeInfo()
|
||||||
|
if runtime.ProviderName != "resp" || runtime.Model != "m2" || runtime.Thinking != "xhigh" {
|
||||||
|
t.Fatalf("agent runtime = %#v", runtime)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestManagerResumeFallsBackWhenRuntimeMissing(t *testing.T) {
|
||||||
|
manager, assistant, _ := testManager(t)
|
||||||
|
|
||||||
|
sess := &Session{
|
||||||
|
ID: GenerateID(),
|
||||||
|
Provider: "ghost",
|
||||||
|
Model: "ghost-model",
|
||||||
|
Thinking: "xhigh",
|
||||||
|
Messages: []llm.Message{{Role: llm.RoleUser, Content: "hi"}},
|
||||||
|
}
|
||||||
|
if err := manager.store.Save(sess); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := manager.Resume(sess.ID); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if got := manager.CurrentProvider(); got != "one" {
|
||||||
|
t.Fatalf("provider = %q, want default one", got)
|
||||||
|
}
|
||||||
|
if got := manager.CurrentModel(); got != "model-a" {
|
||||||
|
t.Fatalf("model = %q, want default model-a", got)
|
||||||
|
}
|
||||||
|
if got := manager.CurrentThinking(); got != "xhigh" {
|
||||||
|
t.Fatalf("thinking = %q, want restored xhigh", got)
|
||||||
|
}
|
||||||
|
runtime := assistant.RuntimeInfo()
|
||||||
|
if runtime.ProviderName != "one" || runtime.Model != "model-a" {
|
||||||
|
t.Fatalf("agent runtime = %#v", runtime)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -22,6 +22,8 @@ type Session struct {
|
|||||||
UpdatedAt time.Time `json:"updated_at"`
|
UpdatedAt time.Time `json:"updated_at"`
|
||||||
Provider string `json:"provider"`
|
Provider string `json:"provider"`
|
||||||
Model string `json:"model"`
|
Model string `json:"model"`
|
||||||
|
APIType string `json:"api_type,omitempty"`
|
||||||
|
Thinking string `json:"thinking,omitempty"`
|
||||||
Messages []llm.Message `json:"messages"`
|
Messages []llm.Message `json:"messages"`
|
||||||
LastUsage *llm.Usage `json:"last_usage,omitempty"`
|
LastUsage *llm.Usage `json:"last_usage,omitempty"`
|
||||||
}
|
}
|
||||||
|
|||||||
+40
-5
@@ -1,6 +1,7 @@
|
|||||||
package tools
|
package tools
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"bytes"
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
@@ -8,6 +9,7 @@ import (
|
|||||||
"os"
|
"os"
|
||||||
"os/exec"
|
"os/exec"
|
||||||
"strings"
|
"strings"
|
||||||
|
"sync"
|
||||||
|
|
||||||
"agentu/pkg/llm"
|
"agentu/pkg/llm"
|
||||||
)
|
)
|
||||||
@@ -46,6 +48,10 @@ type shellRunArgs struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (t *ShellRunTool) Execute(ctx context.Context, raw json.RawMessage) (string, error) {
|
func (t *ShellRunTool) Execute(ctx context.Context, raw json.RawMessage) (string, error) {
|
||||||
|
return t.ExecuteStream(ctx, raw, nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *ShellRunTool) ExecuteStream(ctx context.Context, raw json.RawMessage, emit func(string) error) (string, error) {
|
||||||
var args shellRunArgs
|
var args shellRunArgs
|
||||||
if err := json.Unmarshal(raw, &args); err != nil {
|
if err := json.Unmarshal(raw, &args); err != nil {
|
||||||
return "", fmt.Errorf("parse arguments: %w", err)
|
return "", fmt.Errorf("parse arguments: %w", err)
|
||||||
@@ -66,16 +72,45 @@ func (t *ShellRunTool) Execute(ctx context.Context, raw json.RawMessage) (string
|
|||||||
cmd := exec.CommandContext(ctx, shell, "-lc", args.Command)
|
cmd := exec.CommandContext(ctx, shell, "-lc", args.Command)
|
||||||
cmd.Dir = workdir
|
cmd.Dir = workdir
|
||||||
|
|
||||||
output, err := cmd.CombinedOutput()
|
var mu sync.Mutex
|
||||||
text := FormatResult(string(output))
|
var combined bytes.Buffer
|
||||||
|
writer := &shellStreamWriter{mu: &mu, buf: &combined, emit: emit}
|
||||||
|
cmd.Stdout = writer
|
||||||
|
cmd.Stderr = writer
|
||||||
|
|
||||||
|
runErr := cmd.Run()
|
||||||
|
text := FormatResult(combined.String())
|
||||||
if errors.Is(ctx.Err(), context.DeadlineExceeded) {
|
if errors.Is(ctx.Err(), context.DeadlineExceeded) {
|
||||||
return text, fmt.Errorf("command timed out: %w", ctx.Err())
|
return text, fmt.Errorf("command timed out: %w", ctx.Err())
|
||||||
}
|
}
|
||||||
if err != nil {
|
if runErr != nil {
|
||||||
if text != "" {
|
if text != "" {
|
||||||
return text, fmt.Errorf("command failed: %w", err)
|
return text, fmt.Errorf("command failed: %w", runErr)
|
||||||
}
|
}
|
||||||
return "", fmt.Errorf("command failed: %w", err)
|
return "", fmt.Errorf("command failed: %w", runErr)
|
||||||
}
|
}
|
||||||
return text, nil
|
return text, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// shellStreamWriter captures command output for the final result while
|
||||||
|
// forwarding every write to the streaming callback as it happens.
|
||||||
|
type shellStreamWriter struct {
|
||||||
|
mu *sync.Mutex
|
||||||
|
buf *bytes.Buffer
|
||||||
|
emit func(string) error
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *shellStreamWriter) Write(p []byte) (int, error) {
|
||||||
|
w.mu.Lock()
|
||||||
|
n, writeErr := w.buf.Write(p)
|
||||||
|
w.mu.Unlock()
|
||||||
|
if writeErr != nil {
|
||||||
|
return n, writeErr
|
||||||
|
}
|
||||||
|
if w.emit != nil {
|
||||||
|
if err := w.emit(string(p)); err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return len(p), nil
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,71 @@
|
|||||||
|
package tools
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestShellRunStreamsOutputWhileRunning(t *testing.T) {
|
||||||
|
tool := NewShellRunTool(t.TempDir())
|
||||||
|
var mu sync.Mutex
|
||||||
|
var chunks []string
|
||||||
|
output, err := tool.ExecuteStream(context.Background(), json.RawMessage(`{"command":"printf 'a\\n'; sleep 0.05; printf 'b\\n'"}`), func(chunk string) error {
|
||||||
|
mu.Lock()
|
||||||
|
chunks = append(chunks, chunk)
|
||||||
|
mu.Unlock()
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if !strings.Contains(output, "a") || !strings.Contains(output, "b") {
|
||||||
|
t.Fatalf("output = %q", output)
|
||||||
|
}
|
||||||
|
mu.Lock()
|
||||||
|
joined := strings.Join(chunks, "")
|
||||||
|
mu.Unlock()
|
||||||
|
if len(chunks) == 0 {
|
||||||
|
t.Fatal("expected at least one streamed chunk")
|
||||||
|
}
|
||||||
|
if !strings.Contains(joined, "a") || !strings.Contains(joined, "b") {
|
||||||
|
t.Fatalf("streamed chunks = %q", joined)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestShellRunExecuteCollectsFullOutput(t *testing.T) {
|
||||||
|
tool := NewShellRunTool(t.TempDir())
|
||||||
|
output, err := tool.Execute(context.Background(), json.RawMessage(`{"command":"printf 'hello world'"}`))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(output) != "hello world" {
|
||||||
|
t.Fatalf("output = %q", output)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestShellRunStreamsStderrAndStdout(t *testing.T) {
|
||||||
|
tool := NewShellRunTool(t.TempDir())
|
||||||
|
var mu sync.Mutex
|
||||||
|
var chunks []string
|
||||||
|
output, err := tool.ExecuteStream(context.Background(), json.RawMessage(`{"command":"echo out; echo err >&2"}`), func(chunk string) error {
|
||||||
|
mu.Lock()
|
||||||
|
chunks = append(chunks, chunk)
|
||||||
|
mu.Unlock()
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if !strings.Contains(output, "out") || !strings.Contains(output, "err") {
|
||||||
|
t.Fatalf("output = %q", output)
|
||||||
|
}
|
||||||
|
mu.Lock()
|
||||||
|
joined := strings.Join(chunks, "")
|
||||||
|
mu.Unlock()
|
||||||
|
if !strings.Contains(joined, "out") || !strings.Contains(joined, "err") {
|
||||||
|
t.Fatalf("streamed chunks = %q", joined)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -19,6 +19,15 @@ type Tool interface {
|
|||||||
Execute(ctx context.Context, args json.RawMessage) (string, error)
|
Execute(ctx context.Context, args json.RawMessage) (string, error)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// StreamingTool is optionally implemented by tools that can report partial
|
||||||
|
// output while they execute. emit receives chunks of the tool's output as
|
||||||
|
// they are produced; the returned string is still the complete result that
|
||||||
|
// gets sent back to the model.
|
||||||
|
type StreamingTool interface {
|
||||||
|
Tool
|
||||||
|
ExecuteStream(ctx context.Context, args json.RawMessage, emit func(string) error) (string, error)
|
||||||
|
}
|
||||||
|
|
||||||
// MutatingTool is optionally implemented by tools that modify external state.
|
// MutatingTool is optionally implemented by tools that modify external state.
|
||||||
// It is used by the agent to determine whether tool calls can run in parallel.
|
// It is used by the agent to determine whether tool calls can run in parallel.
|
||||||
type MutatingTool interface {
|
type MutatingTool interface {
|
||||||
|
|||||||
+15
-3
@@ -683,15 +683,27 @@ func (m *model) mergeToolLog(text string) bool {
|
|||||||
if incoming.output == "" {
|
if incoming.output == "" {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
incomingHeader, _, ok := strings.Cut(text, "\n[output]\n")
|
||||||
|
if !ok {
|
||||||
|
return false
|
||||||
|
}
|
||||||
for i := len(m.messages) - 1; i >= 0; i-- {
|
for i := len(m.messages) - 1; i >= 0; i-- {
|
||||||
if m.messages[i].role != roleTool {
|
if m.messages[i].role != roleTool {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
existing := parseToolLog(m.messages[i].content)
|
existing := parseToolLog(m.messages[i].content)
|
||||||
if existing.name == incoming.name && existing.args == incoming.args && existing.output == "" {
|
if existing.name != incoming.name || existing.args != incoming.args {
|
||||||
m.messages[i].content = text
|
continue
|
||||||
return true
|
|
||||||
}
|
}
|
||||||
|
if !strings.Contains(m.messages[i].content, "\n[output]\n") {
|
||||||
|
// Placeholder card from the tool-start log; first chunk fills it.
|
||||||
|
m.messages[i].content = text
|
||||||
|
} else {
|
||||||
|
// Streaming tools send incremental chunks; append the raw delta
|
||||||
|
// so line boundaries are preserved inside the same card.
|
||||||
|
m.messages[i].content += strings.TrimPrefix(text, incomingHeader+"\n[output]\n")
|
||||||
|
}
|
||||||
|
return true
|
||||||
}
|
}
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -395,6 +395,28 @@ func TestToolLogOutputUpdatesExistingToolCard(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestToolLogStreamingChunksAppendToSameCard(t *testing.T) {
|
||||||
|
m := newModel(context.Background(), nil, Options{ModelName: "test-model"})
|
||||||
|
m.messages = append(m.messages, message{role: roleTool, content: `shell_run echo hi`})
|
||||||
|
|
||||||
|
if !m.mergeToolLog("shell_run echo hi\n[output]\none\n") {
|
||||||
|
t.Fatal("expected first mergeToolLog to return true")
|
||||||
|
}
|
||||||
|
if !m.mergeToolLog("shell_run echo hi\n[output]\ntwo\n") {
|
||||||
|
t.Fatal("expected second mergeToolLog to return true")
|
||||||
|
}
|
||||||
|
if len(m.messages) != 1 {
|
||||||
|
t.Fatalf("messages count = %d, want 1", len(m.messages))
|
||||||
|
}
|
||||||
|
plain := stripANSI(m.renderMessages())
|
||||||
|
if !strings.Contains(plain, "one") || !strings.Contains(plain, "two") {
|
||||||
|
t.Fatalf("rendered output missing streamed lines:\n%s", plain)
|
||||||
|
}
|
||||||
|
if strings.Contains(plain, "onetwo") {
|
||||||
|
t.Fatalf("streamed lines were concatenated:\n%s", plain)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
type tuiCompactProvider struct{}
|
type tuiCompactProvider struct{}
|
||||||
|
|
||||||
func (tuiCompactProvider) ChatStream(ctx context.Context, req llm.ChatRequest, emit func(llm.StreamEvent) error) error {
|
func (tuiCompactProvider) ChatStream(ctx context.Context, req llm.ChatRequest, emit func(llm.StreamEvent) error) error {
|
||||||
|
|||||||
+12
-1
@@ -241,7 +241,7 @@ func normalizeProvider(name string, raw rawProviderConfig) ProviderConfig {
|
|||||||
}
|
}
|
||||||
applyModelDefaults(&provider)
|
applyModelDefaults(&provider)
|
||||||
if provider.ThinkingConfigured && provider.ThinkingParam == "" {
|
if provider.ThinkingConfigured && provider.ThinkingParam == "" {
|
||||||
provider.ThinkingParam = "thinking"
|
provider.ThinkingParam = DefaultThinkingParam(provider.APIType)
|
||||||
}
|
}
|
||||||
return provider
|
return provider
|
||||||
}
|
}
|
||||||
@@ -268,6 +268,17 @@ func IsAPIType(value string) bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// DefaultThinkingParam returns the default request field used for the
|
||||||
|
// thinking level. Chat-completions providers historically use a top-level
|
||||||
|
// "thinking" field; the OpenAI Responses API uses "reasoning" with an
|
||||||
|
// {"effort": ...} value.
|
||||||
|
func DefaultThinkingParam(apiType string) string {
|
||||||
|
if apiType == APITypeResponses {
|
||||||
|
return "reasoning"
|
||||||
|
}
|
||||||
|
return "thinking"
|
||||||
|
}
|
||||||
|
|
||||||
func normalizeModelConfigs(rawModels []rawModelConfig) ([]string, map[string]ModelConfig) {
|
func normalizeModelConfigs(rawModels []rawModelConfig) ([]string, map[string]ModelConfig) {
|
||||||
models := make([]string, 0, len(rawModels))
|
models := make([]string, 0, len(rawModels))
|
||||||
configs := make(map[string]ModelConfig, len(rawModels))
|
configs := make(map[string]ModelConfig, len(rawModels))
|
||||||
|
|||||||
@@ -441,3 +441,39 @@ providers:
|
|||||||
t.Fatalf("err = %v", err)
|
t.Fatalf("err = %v", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestLoadDefaultsThinkingParamByAPIType(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
t.Setenv("AGENTU_TEST_KEY", "sk-test")
|
||||||
|
path := filepath.Join(dir, "agentu.yaml")
|
||||||
|
if err := os.WriteFile(path, []byte(`
|
||||||
|
providers:
|
||||||
|
chat:
|
||||||
|
api_type: chat
|
||||||
|
base_url: https://chat.example.com
|
||||||
|
api_key: ${AGENTU_TEST_KEY}
|
||||||
|
models:
|
||||||
|
- id: chat-model
|
||||||
|
thinking: none
|
||||||
|
responses:
|
||||||
|
api_type: responses
|
||||||
|
base_url: https://responses.example.com
|
||||||
|
api_key: ${AGENTU_TEST_KEY}
|
||||||
|
models:
|
||||||
|
- id: responses-model
|
||||||
|
thinking: high
|
||||||
|
`), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := Load(path)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if got := cfg.Providers["chat"].ThinkingParam; got != "thinking" {
|
||||||
|
t.Fatalf("chat thinking param = %q, want %q", got, "thinking")
|
||||||
|
}
|
||||||
|
if got := cfg.Providers["responses"].ThinkingParam; got != "reasoning" {
|
||||||
|
t.Fatalf("responses thinking param = %q, want %q", got, "reasoning")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+22
-1
@@ -124,6 +124,9 @@ func buildResponsesRequest(req ChatRequest) map[string]any {
|
|||||||
"input": buildResponsesInput(req.Messages),
|
"input": buildResponsesInput(req.Messages),
|
||||||
"stream": true,
|
"stream": true,
|
||||||
}
|
}
|
||||||
|
if instructions := responsesInstructions(req.Messages); instructions != "" {
|
||||||
|
payload["instructions"] = instructions
|
||||||
|
}
|
||||||
if len(req.Tools) > 0 {
|
if len(req.Tools) > 0 {
|
||||||
payload["tools"] = buildResponsesTools(req.Tools)
|
payload["tools"] = buildResponsesTools(req.Tools)
|
||||||
}
|
}
|
||||||
@@ -153,7 +156,9 @@ func buildResponsesInput(messages []Message) []any {
|
|||||||
for _, msg := range messages {
|
for _, msg := range messages {
|
||||||
switch msg.Role {
|
switch msg.Role {
|
||||||
case RoleSystem:
|
case RoleSystem:
|
||||||
items = append(items, responsesMessageInput{Role: "system", Content: msg.Content})
|
// System instructions are sent as the top-level "instructions"
|
||||||
|
// field and must not be duplicated inside the input array.
|
||||||
|
continue
|
||||||
case RoleUser:
|
case RoleUser:
|
||||||
items = append(items, responsesMessageInput{Role: "user", Content: msg.Content})
|
items = append(items, responsesMessageInput{Role: "user", Content: msg.Content})
|
||||||
case RoleAssistant:
|
case RoleAssistant:
|
||||||
@@ -179,6 +184,22 @@ func buildResponsesInput(messages []Message) []any {
|
|||||||
return items
|
return items
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// responsesInstructions extracts system messages into a single instructions
|
||||||
|
// string, which is the OpenAI Responses API way to pass developer/system
|
||||||
|
// guidance (and the most prompt-cache friendly).
|
||||||
|
func responsesInstructions(messages []Message) string {
|
||||||
|
var instructions []string
|
||||||
|
for _, msg := range messages {
|
||||||
|
if msg.Role != RoleSystem {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if text := strings.TrimSpace(msg.Content); text != "" {
|
||||||
|
instructions = append(instructions, text)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return strings.Join(instructions, "\n\n")
|
||||||
|
}
|
||||||
|
|
||||||
func buildResponsesTools(tools []Tool) []responsesTool {
|
func buildResponsesTools(tools []Tool) []responsesTool {
|
||||||
out := make([]responsesTool, 0, len(tools))
|
out := make([]responsesTool, 0, len(tools))
|
||||||
for _, tool := range tools {
|
for _, tool := range tools {
|
||||||
|
|||||||
@@ -59,12 +59,15 @@ func TestBuildResponsesRequest(t *testing.T) {
|
|||||||
if raw["model"] != "gpt-test" || raw["stream"] != true {
|
if raw["model"] != "gpt-test" || raw["stream"] != true {
|
||||||
t.Fatalf("model/stream = %#v / %#v", raw["model"], raw["stream"])
|
t.Fatalf("model/stream = %#v / %#v", raw["model"], raw["stream"])
|
||||||
}
|
}
|
||||||
|
if raw["instructions"] != "be concise" {
|
||||||
|
t.Fatalf("instructions = %#v", raw["instructions"])
|
||||||
|
}
|
||||||
if raw["tool_choice"] != "auto" {
|
if raw["tool_choice"] != "auto" {
|
||||||
t.Fatalf("tool_choice = %#v", raw["tool_choice"])
|
t.Fatalf("tool_choice = %#v", raw["tool_choice"])
|
||||||
}
|
}
|
||||||
|
|
||||||
input, ok := raw["input"].([]any)
|
input, ok := raw["input"].([]any)
|
||||||
if !ok || len(input) != 6 {
|
if !ok || len(input) != 5 {
|
||||||
t.Fatalf("input = %#v", raw["input"])
|
t.Fatalf("input = %#v", raw["input"])
|
||||||
}
|
}
|
||||||
checkInputItem := func(index int, wantType, wantKey, wantValue string) {
|
checkInputItem := func(index int, wantType, wantKey, wantValue string) {
|
||||||
@@ -80,16 +83,15 @@ func TestBuildResponsesRequest(t *testing.T) {
|
|||||||
t.Fatalf("input[%d].type = %#v, want %q", index, item["type"], wantType)
|
t.Fatalf("input[%d].type = %#v, want %q", index, item["type"], wantType)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
checkInputItem(0, "", "role", "system")
|
checkInputItem(0, "", "role", "user")
|
||||||
checkInputItem(1, "", "role", "user")
|
checkInputItem(1, "", "role", "assistant")
|
||||||
checkInputItem(2, "", "role", "assistant")
|
checkInputItem(2, "function_call", "call_id", "call_1")
|
||||||
checkInputItem(3, "function_call", "call_id", "call_1")
|
checkInputItem(3, "function_call_output", "call_id", "call_1")
|
||||||
checkInputItem(4, "function_call_output", "call_id", "call_1")
|
checkInputItem(4, "", "role", "assistant")
|
||||||
checkInputItem(5, "", "role", "assistant")
|
if item := input[2].(map[string]any); item["name"] != "file_read" || item["arguments"] != `{"path":"README.md"}` {
|
||||||
if item := input[3].(map[string]any); item["name"] != "file_read" || item["arguments"] != `{"path":"README.md"}` {
|
|
||||||
t.Fatalf("function_call item = %#v", item)
|
t.Fatalf("function_call item = %#v", item)
|
||||||
}
|
}
|
||||||
if item := input[4].(map[string]any); item["output"] != "contents" {
|
if item := input[3].(map[string]any); item["output"] != "contents" {
|
||||||
t.Fatalf("function_call_output item = %#v", item)
|
t.Fatalf("function_call_output item = %#v", item)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -117,6 +119,36 @@ func TestBuildResponsesRequest(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestResponsesInstructionsCombinesSystemMessages(t *testing.T) {
|
||||||
|
payload := buildResponsesRequest(ChatRequest{
|
||||||
|
Model: "gpt-test",
|
||||||
|
Messages: []Message{
|
||||||
|
{Role: RoleSystem, Content: "first instruction"},
|
||||||
|
{Role: RoleUser, Content: "hi"},
|
||||||
|
{Role: RoleSystem, Content: "second instruction"},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
data, err := json.Marshal(payload)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
var raw map[string]any
|
||||||
|
if err := json.Unmarshal(data, &raw); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if raw["instructions"] != "first instruction\n\nsecond instruction" {
|
||||||
|
t.Fatalf("instructions = %#v", raw["instructions"])
|
||||||
|
}
|
||||||
|
input, ok := raw["input"].([]any)
|
||||||
|
if !ok || len(input) != 1 {
|
||||||
|
t.Fatalf("input = %#v", raw["input"])
|
||||||
|
}
|
||||||
|
item, ok := input[0].(map[string]any)
|
||||||
|
if !ok || item["role"] != "user" {
|
||||||
|
t.Fatalf("input[0] = %#v", input[0])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestOpenAIResponsesClientStreamsContentAndToolCalls(t *testing.T) {
|
func TestOpenAIResponsesClientStreamsContentAndToolCalls(t *testing.T) {
|
||||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
if r.URL.Path != "/v1/responses" {
|
if r.URL.Path != "/v1/responses" {
|
||||||
|
|||||||
Reference in New Issue
Block a user