Files
robinhood-agentic-mcp/market/market_test.go
T
s1d3sw1ped a6bf8632ce fix: parse MCP result payloads and drop place idempotency_key
Typed result structs replace empty envelopes. Equity place sends
ref_id only so live additionalProperties:false schemas accept the call.
2026-09-01 14:18:39 -05:00

258 lines
6.9 KiB
Go

package market_test
import (
"context"
"encoding/json"
"sort"
"testing"
"time"
"github.com/google/go-cmp/cmp"
"s1d3sw1ped/robinhood-agentic-mcp/client"
"s1d3sw1ped/robinhood-agentic-mcp/internal/rhntest"
"s1d3sw1ped/robinhood-agentic-mcp/market"
)
func TestMarket_toolNames(t *testing.T) {
t.Parallel()
start := time.Date(2026, 8, 18, 13, 30, 0, 0, time.UTC)
end := time.Date(2026, 8, 18, 20, 0, 0, 0, time.UTC)
tests := []struct {
name string
call func(*market.Client) error
wantName string
wantArgs map[string]any
}{
{
name: "Indexes",
call: func(c *market.Client) error {
_, err := c.Indexes(context.Background(), market.IndexesRequest{Symbols: "SPX,NDX"})
return err
},
wantName: "get_indexes",
wantArgs: map[string]any{"symbols": "SPX,NDX"},
},
{
name: "IndexQuotes",
call: func(c *market.Client) error {
_, err := c.IndexQuotes(context.Background(), market.IndexQuotesRequest{
InstrumentIDs: []string{"idx-spx", "idx-ndx"},
})
return err
},
wantName: "get_index_quotes",
wantArgs: map[string]any{"instrument_ids": []string{"idx-spx", "idx-ndx"}},
},
{
name: "IndexHistoricals",
call: func(c *market.Client) error {
_, err := c.IndexHistoricals(context.Background(), market.IndexHistoricalsRequest{
InstrumentIDs: []string{"idx-spx"},
StartTime: start,
EndTime: end,
Interval: "day",
})
return err
},
wantName: "get_index_historicals",
wantArgs: map[string]any{
"instrument_ids": []string{"idx-spx"},
"start_time": "2026-08-18T13:30:00Z",
"end_time": "2026-08-18T20:00:00Z",
"interval": "day",
},
},
{
name: "Financials",
call: func(c *market.Client) error {
_, err := c.Financials(context.Background(), market.FinancialsRequest{
Symbols: []string{"MU", "AAPL"},
Period: "quarterly",
Limit: 8,
})
return err
},
wantName: "get_financials",
wantArgs: map[string]any{
"symbols": []string{"MU", "AAPL"},
"period": "quarterly",
"limit": 8,
},
},
{
name: "EarningsResults",
call: func(c *market.Client) error {
_, err := c.EarningsResults(context.Background(), market.EarningsResultsRequest{Symbol: "MU"})
return err
},
wantName: "get_earnings_results",
wantArgs: map[string]any{"symbol": "MU"},
},
{
name: "EarningsCalendar",
call: func(c *market.Client) error {
_, err := c.EarningsCalendar(context.Background(), market.EarningsCalendarRequest{
StartDate: "2026-08-18",
Days: 7,
Filter: "high_market_cap",
})
return err
},
wantName: "get_earnings_calendar",
wantArgs: map[string]any{
"start_date": "2026-08-18",
"days": 7,
"filter": "high_market_cap",
},
},
{
name: "SECFilingIndex",
call: func(c *market.Client) error {
_, err := c.SECFilingIndex(context.Background(), market.SECFilingIndexRequest{
Symbol: "MU",
FormType: []string{"10-K", "10-Q"},
Since: "2026-01-01",
Until: "2026-08-18",
Cursor: "c1",
})
return err
},
wantName: "get_sec_filing_index",
wantArgs: map[string]any{
"symbol": "MU",
"form_type": []string{"10-K", "10-Q"},
"since": "2026-01-01",
"until": "2026-08-18",
"cursor": "c1",
},
},
{
name: "SECFiling",
call: func(c *market.Client) error {
_, err := c.SECFiling(context.Background(), market.SECFilingRequest{
FilingID: "f-1",
Section: "item1",
})
return err
},
wantName: "get_sec_filing",
wantArgs: map[string]any{"filing_id": "f-1", "section": "item1"},
},
{
name: "SECFilingFacts",
call: func(c *market.Client) error {
_, err := c.SECFilingFacts(context.Background(), market.SECFilingFactsRequest{
FilingIDs: []string{"f-1", "f-2"},
Concepts: []string{"NetIncomeLoss", "Revenues"},
})
return err
},
wantName: "get_sec_filing_facts",
wantArgs: map[string]any{
"filing_ids": []string{"f-1", "f-2"},
"concepts": []string{"NetIncomeLoss", "Revenues"},
},
},
{
name: "SECFilingFactsCatalog",
call: func(c *market.Client) error {
_, err := c.SECFilingFactsCatalog(context.Background(), market.SECFilingFactsCatalogRequest{
FilingID: "f-1",
ConceptContains: "Debt",
AxisNameIn: []string{"LegalEntityAxis"},
Offset: 10,
})
return err
},
wantName: "get_sec_filing_facts_catalog",
wantArgs: map[string]any{
"filing_id": "f-1",
"concept_contains": "Debt",
"axis_name_in": []string{"LegalEntityAxis"},
"offset": 10,
},
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
var gotName string
var gotArgs map[string]any
c := market.New(client.Func(func(ctx context.Context, name string, args map[string]any) (json.RawMessage, error) {
gotName, gotArgs = name, args
return json.RawMessage(`{}`), nil
}))
if err := tc.call(c); err != nil {
t.Fatal(err)
}
if gotName != tc.wantName {
t.Fatalf("%s %+v", gotName, gotArgs)
}
if diff := cmp.Diff(tc.wantArgs, gotArgs); diff != "" {
t.Fatal(diff)
}
if tc.name == "IndexHistoricals" {
if gotArgs["interval"] != "day" {
t.Fatalf("interval %v", gotArgs["interval"])
}
}
})
}
}
func TestTools(t *testing.T) {
t.Parallel()
want := []string{
"get_earnings_calendar",
"get_earnings_results",
"get_financials",
"get_index_historicals",
"get_index_quotes",
"get_indexes",
"get_sec_filing",
"get_sec_filing_facts",
"get_sec_filing_facts_catalog",
"get_sec_filing_index",
}
got := append([]string(nil), market.Tools()...)
sort.Strings(got)
if diff := cmp.Diff(want, got); diff != "" {
t.Fatal(diff)
}
}
func TestEarningsResults_rhntest(t *testing.T) {
t.Parallel()
s := rhntest.New(t)
s.Set("get_earnings_results", json.RawMessage(`{"next_report_date":"2026-10-15","report_date":"2026-07-15"}`))
c := market.New(&client.Client{URL: s.URL})
got, err := c.EarningsResults(context.Background(), market.EarningsResultsRequest{Symbol: "MU"})
if err != nil {
t.Fatal(err)
}
if got.NextReportDate != "2026-10-15" || got.ReportDate != "2026-07-15" {
t.Fatalf("%+v", got)
}
if s.LastName() != "get_earnings_results" {
t.Fatalf("%s", s.LastName())
}
want := map[string]any{"symbol": "MU"}
if diff := cmp.Diff(want, s.LastArgs()); diff != "" {
t.Fatal(diff)
}
}
func TestEarningsResults_liveEnvelope(t *testing.T) {
t.Parallel()
c := market.New(client.Func(func(ctx context.Context, name string, args map[string]any) (json.RawMessage, error) {
return json.RawMessage(`{"data":{"results":[{"report":{"date":"2026-07-15"}}]}}`), nil
}))
got, err := c.EarningsResults(context.Background(), market.EarningsResultsRequest{Symbol: "MU"})
if err != nil {
t.Fatal(err)
}
if got.ReportDate != "2026-07-15" || got.NextReportDate != "2026-07-15" {
t.Fatalf("%+v", got)
}
}