From dded81a44382f30508ac18fc78b996095b1d511e Mon Sep 17 00:00:00 2001 From: PratikDhanave Date: Fri, 24 Jul 2026 06:58:04 +0530 Subject: [PATCH] Make Foundry MemoryProvider store message filters configurable Add StoreInputRequestFilter and StoreInputResponseFilter to MemoryProviderConfig and wire them into the underlying ContextProviderConfig. Previously the request store filter was hardcoded to ExternalOnly and the response store filter was left unset, so callers could not customize which request or response messages are persisted as memories. Defaults are preserved to match .NET/Python parity: request messages default to ExternalOnly and response messages default to PassThrough. --- provider/foundryprovider/memory.go | 21 ++++++- provider/foundryprovider/memory_test.go | 73 +++++++++++++++++++++++++ 2 files changed, 91 insertions(+), 3 deletions(-) diff --git a/provider/foundryprovider/memory.go b/provider/foundryprovider/memory.go index 5633fc21..f25afc1a 100644 --- a/provider/foundryprovider/memory.go +++ b/provider/foundryprovider/memory.go @@ -40,6 +40,14 @@ type MemoryProviderConfig struct { // default is [messagefilter.ExternalOnly]. SearchInputFilter messagefilter.Filter + // StoreInputRequestFilter filters request messages before they are stored as + // memories. The default is [messagefilter.ExternalOnly]. + StoreInputRequestFilter messagefilter.Filter + + // StoreInputResponseFilter filters response messages before they are stored as + // memories. The default is [messagefilter.PassThrough]. + StoreInputResponseFilter messagefilter.Filter + // UpdateDelay controls Foundry memory extraction delay in seconds. The default is 0, // which submits memory updates immediately. UpdateDelay int32 @@ -98,10 +106,17 @@ func newMemoryProvider(client *azaiprojects.MemoryStoresClient, memoryStoreName if config.SearchInputFilter == nil { config.SearchInputFilter = messagefilter.ExternalOnly } + if config.StoreInputRequestFilter == nil { + config.StoreInputRequestFilter = messagefilter.ExternalOnly + } + if config.StoreInputResponseFilter == nil { + config.StoreInputResponseFilter = messagefilter.PassThrough + } providerConfig := agent.ContextProviderConfig{ - ProvideInputMessageFilter: config.SearchInputFilter, - SourceID: defaultSourceID, - StoreInputRequestMessageFilter: messagefilter.ExternalOnly, + ProvideInputMessageFilter: config.SearchInputFilter, + SourceID: defaultSourceID, + StoreInputRequestMessageFilter: config.StoreInputRequestFilter, + StoreInputResponseMessageFilter: config.StoreInputResponseFilter, } p := &MemoryProvider{ client: client, diff --git a/provider/foundryprovider/memory_test.go b/provider/foundryprovider/memory_test.go index 96405e2c..70d0ec60 100644 --- a/provider/foundryprovider/memory_test.go +++ b/provider/foundryprovider/memory_test.go @@ -86,6 +86,79 @@ func TestNewMemoryProviderUsesCustomSearchInputFilter(t *testing.T) { } } +func TestNewMemoryProviderUsesCustomStoreFilters(t *testing.T) { + transport := &recordingTransport{} + transport.handle = func(req *http.Request, _ string) (*http.Response, error) { + resp := jsonResponse(req, http.StatusAccepted, `{"update_id":"update_1","status":"queued"}`) + resp.Header.Set("Operation-Location", validEndpoint+"/memory_stores/memory/updates/update_1?api-version=v1") + return resp, nil + } + requestCalled := false + responseCalled := false + requestFilter := func(_ context.Context, messages []*message.Message) ([]*message.Message, error) { + requestCalled = true + return messages, nil + } + responseFilter := func(_ context.Context, messages []*message.Message) ([]*message.Message, error) { + responseCalled = true + return messages, nil + } + provider := foundryprovider.NewMemoryProvider(validEndpoint, validCredential, "memory", validScope, foundryprovider.MemoryProviderConfig{ + ClientOptions: azcore.ClientOptions{Transport: transport}, + StoreInputRequestFilter: requestFilter, + StoreInputResponseFilter: responseFilter, + }) + + err := provider.Invoked(t.Context(), agent.InvokedContext{ + RequestMessages: []*message.Message{message.NewText("remember me")}, + ResponseMessages: []*message.Message{{Role: message.RoleAssistant, Contents: message.Contents{&message.TextContent{Text: "assistant text"}}}}, + }) + if err != nil { + t.Fatalf("Invoked error = %v", err) + } + if !requestCalled { + t.Fatal("custom store request filter was not called") + } + if !responseCalled { + t.Fatal("custom store response filter was not called") + } +} + +func TestNewMemoryProviderStoreFiltersDefaultToExternalOnlyRequestAndPassThroughResponse(t *testing.T) { + transport := &recordingTransport{} + transport.handle = func(req *http.Request, _ string) (*http.Response, error) { + resp := jsonResponse(req, http.StatusAccepted, `{"update_id":"update_1","status":"queued"}`) + resp.Header.Set("Operation-Location", validEndpoint+"/memory_stores/memory/updates/update_1?api-version=v1") + return resp, nil + } + provider := foundryprovider.NewMemoryProvider(validEndpoint, validCredential, "memory", validScope, foundryprovider.MemoryProviderConfig{ + ClientOptions: azcore.ClientOptions{Transport: transport}, + }) + + // Request message with a non-external source is dropped by the default + // ExternalOnly request filter; response message with a non-external source is + // kept by the default PassThrough response filter. + err := provider.Invoked(t.Context(), agent.InvokedContext{ + RequestMessages: []*message.Message{{Role: message.RoleUser, Source: message.Source{Type: agent.SourceTypeContextProvider}, Contents: message.Contents{&message.TextContent{Text: "internal request"}}}}, + ResponseMessages: []*message.Message{{Role: message.RoleAssistant, Source: message.Source{Type: agent.SourceTypeContextProvider}, Contents: message.Contents{&message.TextContent{Text: "internal response"}}}}, + }) + if err != nil { + t.Fatalf("Invoked error = %v", err) + } + + requests := transport.Requests() + if len(requests) != 1 { + t.Fatalf("request count = %d, want 1", len(requests)) + } + items, ok := jsonMap(t, requests[0].Body)["items"].([]any) + if !ok || len(items) != 1 { + t.Fatalf("items = %#v, want only the response message", items) + } + if items[0].(map[string]any)["role"] != "assistant" { + t.Fatalf("item role = %#v, want assistant", items[0]) + } +} + func TestMemoryProviderPanicsWhenScopeIsEmptyOnUse(t *testing.T) { provider := foundryprovider.NewMemoryProvider(validEndpoint, validCredential, "memory", func(*agent.Session) string { return " " }, foundryprovider.MemoryProviderConfig{}) assertPanics(t, func() {