diff --git a/typesense/multi_search.go b/typesense/multi_search.go index 19e1cf7c..1a20b4cb 100644 --- a/typesense/multi_search.go +++ b/typesense/multi_search.go @@ -4,13 +4,16 @@ import ( "bytes" "context" "encoding/json" + "errors" "io" + "net/http" "github.com/typesense/typesense-go/v4/typesense/api" ) type MultiSearchInterface interface { Perform(ctx context.Context, commonSearchParams *api.MultiSearchParams, searchParams api.MultiSearchSearchesParameter) (*api.MultiSearchResult, error) + PerformUnion(ctx context.Context, commonSearchParams *api.MultiSearchParams, searchParams api.MultiSearchSearchesParameter) (*api.SearchResult, error) PerformWithContentType(ctx context.Context, commonSearchParams *api.MultiSearchParams, searchParams api.MultiSearchSearchesParameter, contentType string) (*api.MultiSearchResponse, error) } @@ -19,6 +22,9 @@ type multiSearch struct { } func (m *multiSearch) Perform(ctx context.Context, commonSearchParams *api.MultiSearchParams, searchParams api.MultiSearchSearchesParameter) (*api.MultiSearchResult, error) { + if searchParams.Union != nil && *searchParams.Union { + return nil, errors.New("union must be false for Perform; use PerformUnion") + } response, err := m.apiClient.MultiSearchWithResponse(ctx, commonSearchParams, api.MultiSearchJSONRequestBody(searchParams)) if err != nil { return nil, err @@ -29,6 +35,34 @@ func (m *multiSearch) Perform(ctx context.Context, commonSearchParams *api.Multi return response.JSON200, nil } +func (m *multiSearch) PerformUnion(ctx context.Context, commonSearchParams *api.MultiSearchParams, searchParams api.MultiSearchSearchesParameter) (*api.SearchResult, error) { + if searchParams.Union != nil && !*searchParams.Union { + return nil, errors.New("union must be true for PerformUnion") + } + union := true + searchParams.Union = &union + + httpResp, err := m.apiClient.MultiSearch(ctx, commonSearchParams, api.MultiSearchJSONRequestBody(searchParams)) + if err != nil { + return nil, err + } + defer httpResp.Body.Close() + + responseBody, err := io.ReadAll(httpResp.Body) + if err != nil { + return nil, err + } + if httpResp.StatusCode < http.StatusOK || httpResp.StatusCode >= http.StatusMultipleChoices { + return nil, &HTTPError{Status: httpResp.StatusCode, Body: responseBody} + } + + var result api.SearchResult + if err := json.Unmarshal(responseBody, &result); err != nil { + return nil, err + } + return &result, nil +} + func (m *multiSearch) PerformWithContentType(ctx context.Context, commonSearchParams *api.MultiSearchParams, searchParams api.MultiSearchSearchesParameter, contentType string) (*api.MultiSearchResponse, error) { body := api.MultiSearchJSONRequestBody(searchParams) var requestReader io.Reader diff --git a/typesense/multi_search_test.go b/typesense/multi_search_test.go index 27638b11..8e28eb1f 100644 --- a/typesense/multi_search_test.go +++ b/typesense/multi_search_test.go @@ -1,14 +1,14 @@ package typesense import ( + "bytes" "context" "encoding/json" "errors" + "io" "net/http" "testing" - "bytes" - "github.com/stretchr/testify/assert" "github.com/typesense/typesense-go/v4/typesense/api" "github.com/typesense/typesense-go/v4/typesense/api/pointer" @@ -206,6 +206,78 @@ func TestMultiSearchResultDeserialization(t *testing.T) { assert.Equal(t, expected, result) } +func TestMultiSearchPerformUnion(t *testing.T) { + expectedParams := newMultiSearchParams() + expectedBody := newMultiSearchBodyParams() + union := true + expectedBody.Union = &union + + expectedResult := &api.SearchResult{ + Found: pointer.Int(2), + Hits: &[]api.SearchResultHit{ + { + Document: &map[string]interface{}{ + "id": "124", + "company_name": "Stark Industries", + "num_employees": float64(5215), + "country": "USA", + }, + SearchIndex: pointer.Int(1), + }, + }, + OutOf: pointer.Int(179203), + Page: pointer.Int(0), + SearchCutoff: pointer.False(), + SearchTimeMs: pointer.Int(8), + UnionRequestParams: &[]api.SearchRequestParams{ + { + CollectionName: "events", + Q: "iran", + PerPage: 10, + }, + { + CollectionName: "events", + Q: "iran", + PerPage: 10, + }, + }, + } + + expectedResponseBytes, err := json.Marshal(expectedResult) + assert.Nil(t, err) + + ctrl := gomock.NewController(t) + defer ctrl.Finish() + mockAPIClient := mocks.NewMockAPIClientInterface(ctrl) + mockAPIClient.EXPECT(). + MultiSearch(gomock.Not(gomock.Nil()), expectedParams, api.MultiSearchJSONRequestBody(expectedBody)). + Return(&http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewReader(expectedResponseBytes)), + }, nil).Times(1) + + client := NewClient(WithAPIClient(mockAPIClient)) + params := newMultiSearchParams() + body := newMultiSearchBodyParams() + result, err := client.MultiSearch.PerformUnion(context.Background(), params, body) + + assert.Nil(t, err) + assert.Equal(t, expectedResult, result) +} + +func TestMultiSearchPerformUnionRejectsUnionFalse(t *testing.T) { + ctrl := gomock.NewController(t) + defer ctrl.Finish() + client := NewClient(WithAPIClient(mocks.NewMockAPIClientInterface(ctrl))) + params := newMultiSearchParams() + body := newMultiSearchBodyParams() + union := false + body.Union = &union + + _, err := client.MultiSearch.PerformUnion(context.Background(), params, body) + assert.NotNil(t, err) +} + func TestMultiSearch(t *testing.T) { expectedParams := newMultiSearchParams() expectedResult := newMultiSearchResult() diff --git a/typesense/test/multi_search_test.go b/typesense/test/multi_search_test.go index 90912f4e..b9d5013e 100644 --- a/typesense/test/multi_search_test.go +++ b/typesense/test/multi_search_test.go @@ -346,3 +346,51 @@ func TestMultiSearchWithStopwords(t *testing.T) { // Check second result require.Equal(t, 0, len(*result.Results[1].Hits), "Number of docs in second result did not equal") } + +func TestMultiSearchUnion(t *testing.T) { + collectionName1 := createNewCollection(t, "companies") + collectionName2 := createNewCollection(t, "companies") + + createDocument(t, collectionName1, newDocument("201", withCompanyName("Stark Alpha"), withNumEmployees(10))) + createDocument(t, collectionName2, newDocument("301", withCompanyName("Stark Beta"), withNumEmployees(20))) + + params := &api.MultiSearchParams{ + Page: pointer.Int(1), + PerPage: pointer.Int(10), + } + + searches := api.MultiSearchSearchesParameter{ + Searches: []api.MultiSearchCollectionParameters{ + { + Collection: pointer.String(collectionName1), + Q: pointer.String("*"), + QueryBy: pointer.String("company_name"), + FilterBy: pointer.String("company_name:Stark Alpha"), + }, + { + Collection: pointer.String(collectionName2), + Q: pointer.String("*"), + QueryBy: pointer.String("company_name"), + FilterBy: pointer.String("company_name:Stark Beta"), + }, + }, + } + + result, err := typesenseClient.MultiSearch.PerformUnion(context.Background(), params, searches) + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, result.Hits) + require.Equal(t, 2, len(*result.Hits)) + require.NotNil(t, result.UnionRequestParams) + require.Equal(t, 2, len(*result.UnionRequestParams)) + + foundNames := map[string]bool{} + for _, hit := range *result.Hits { + require.NotNil(t, hit.Document) + if name, ok := (*hit.Document)["company_name"].(string); ok { + foundNames[name] = true + } + } + require.True(t, foundNames["Stark Alpha"]) + require.True(t, foundNames["Stark Beta"]) +}