Skip to content
22 changes: 16 additions & 6 deletions provider/foundryprovider/memory.go
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,14 @@ type MemoryProviderConfig struct {
// default is [messagefilter.ExternalOnly].
SearchInputFilter messagefilter.Filter

// StorageInputRequestMessageFilter filters request messages before they are stored as
// memories. The default is [messagefilter.ExternalOnly].
StorageInputRequestMessageFilter messagefilter.Filter

// StorageInputResponseMessageFilter filters response messages before they are stored as
// memories. The default is [messagefilter.PassThrough].
StorageInputResponseMessageFilter messagefilter.Filter

// UpdateDelay controls Foundry memory extraction delay in seconds. The default is 0,
// which submits memory updates immediately.
UpdateDelay int32
Expand Down Expand Up @@ -97,13 +105,15 @@ func newMemoryProvider(client *azaiprojects.MemoryStoresClient, memoryStoreName
if config.MaxMemories == 0 {
config.MaxMemories = defaultMaxMemories
}
if config.SearchInputFilter == nil {
config.SearchInputFilter = messagefilter.ExternalOnly
}
// All three filters are intentionally left nil when unset: ContextProviderConfig
// already defaults a nil provide/store-request filter to messagefilter.ExternalOnly
// and a nil store-response filter to messagefilter.PassThrough, so re-setting them
// here would be redundant.
providerConfig := agent.ContextProviderConfig{
ProvideInputMessageFilter: config.SearchInputFilter,
SourceID: defaultSourceID,
StoreInputRequestMessageFilter: messagefilter.ExternalOnly,
ProvideInputMessageFilter: config.SearchInputFilter,
SourceID: defaultSourceID,
StoreInputRequestMessageFilter: config.StorageInputRequestMessageFilter,
StoreInputResponseMessageFilter: config.StorageInputResponseMessageFilter,
}
p := &MemoryProvider{
client: client,
Expand Down
73 changes: 73 additions & 0 deletions provider/foundryprovider/memory_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -87,6 +87,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},
StorageInputRequestMessageFilter: requestFilter,
StorageInputResponseMessageFilter: 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() {
Expand Down
Loading