Merge pull request #273 from iceljc/features/add-truncate-message

add truncate message
This commit is contained in:
Haiping 2024-01-28 19:37:45 -06:00 committed by GitHub
commit 0c066d6b4a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 154 additions and 4 deletions

View file

@ -13,6 +13,7 @@ public interface IConversationService
Task<Conversation> UpdateConversationTitle(string id, string title);
Task<List<Conversation>> GetLastConversations();
Task<bool> DeleteConversation(string id);
Task<bool> TruncateConversation(string conversationId, string messageId);
/// <summary>
/// Send message to LLM

View file

@ -0,0 +1,6 @@
namespace BotSharp.Abstraction.Conversations.Models;
public class TruncateMessageRequest
{
public string? TruncateMessageId { get; set; }
}

View file

@ -1,6 +1,6 @@
namespace BotSharp.Abstraction.Models;
public class MessageConfig
public class MessageConfig : TruncateMessageRequest
{
/// <summary>
/// Completion Provider

View file

@ -46,6 +46,7 @@ public interface IBotSharpRepository
PagedItems<Conversation> GetConversations(ConversationFilter filter);
void UpdateConversationTitle(string conversationId, string title);
List<Conversation> GetLastConversations();
bool TruncateConversation(string conversationId, string messageId);
#endregion
#region Execution Log

View file

@ -0,0 +1,13 @@
using BotSharp.Abstraction.Repositories;
namespace BotSharp.Core.Conversations.Services;
public partial class ConversationService : IConversationService
{
public async Task<bool> TruncateConversation(string conversationId, string messageId)
{
var db = _services.GetRequiredService<IBotSharpRepository>();
var isSaved = db.TruncateConversation(conversationId, messageId);
return await Task.FromResult(isSaved);
}
}

View file

@ -178,6 +178,11 @@ public class BotSharpDbContext : Database, IBotSharpRepository
{
throw new NotImplementedException();
}
public bool TruncateConversation(string conversationId, string messageId)
{
throw new NotImplementedException();
}
#endregion

View file

@ -1,5 +1,6 @@
using BotSharp.Abstraction.Repositories.Filters;
using BotSharp.Abstraction.Repositories.Models;
using System.Globalization;
using System.IO;
namespace BotSharp.Core.Repository
@ -256,6 +257,35 @@ namespace BotSharp.Core.Repository
}
public bool TruncateConversation(string conversationId, string messageId)
{
if (string.IsNullOrEmpty(conversationId) || string.IsNullOrEmpty(messageId)) return false;
var dialogs = new List<DialogElement>();
var convDir = FindConversationDirectory(conversationId);
if (string.IsNullOrEmpty(convDir)) return false;
var dialogDir = Path.Combine(convDir, DIALOG_FILE);
dialogs = CollectDialogElements(dialogDir);
if (dialogs.IsNullOrEmpty()) return false;
var foundIdx = dialogs.FindIndex(x => x.MetaData?.MessageId == messageId);
if (foundIdx < 0) return false;
// Handle truncated dialogs
var isSaved = HandleTruncatedDialogs(dialogDir, dialogs, foundIdx);
if (!isSaved) return false;
// Handle truncated states
var refTime = dialogs.ElementAt(foundIdx).MetaData.CreateTime;
var stateDir = Path.Combine(convDir, STATE_FILE);
var states = CollectConversationStates(stateDir);
isSaved = HandleTruncatedStates(stateDir, states, refTime);
return isSaved;
}
#region Private methods
private string? FindConversationDirectory(string conversationId)
{
@ -304,8 +334,9 @@ namespace BotSharp.Core.Repository
foreach (var element in dialogs)
{
var meta = element.MetaData;
var createTime = meta.CreateTime.ToString("MM/dd/yyyy hh:mm:ss.fff tt", CultureInfo.InvariantCulture);
var source = meta.FunctionName ?? meta.SenderId;
var metaStr = $"{meta.CreateTime}|{meta.Role}|{meta.AgentId}|{meta.MessageId}|{source}";
var metaStr = $"{createTime}|{meta.Role}|{meta.AgentId}|{meta.MessageId}|{source}";
dialogTexts.Add(metaStr);
var content = $" - {element.Content}";
dialogTexts.Add(content);
@ -325,6 +356,49 @@ namespace BotSharp.Core.Repository
states = JsonSerializer.Deserialize<List<StateKeyValue>>(stateStr, _options);
return states ?? new List<StateKeyValue>();
}
private bool HandleTruncatedDialogs(string dialogDir, List<DialogElement> dialogs, int foundIdx)
{
var truncatedDialogs = dialogs.Where((x, idx) => idx < foundIdx).ToList();
var isSaved = SaveTruncatedDialogs(dialogDir, truncatedDialogs);
return isSaved;
}
private bool HandleTruncatedStates(string stateDir, List<StateKeyValue> states, DateTime refTime)
{
var truncatedStates = new List<StateKeyValue>();
foreach (var state in states)
{
var values = state.Values.Where(x => x.UpdateTime < refTime).ToList();
if (values.Count == 0) continue;
state.Values = values;
truncatedStates.Add(state);
}
var isSaved = SaveTruncatedStates(stateDir, truncatedStates);
return isSaved;
}
private bool SaveTruncatedDialogs(string dialogDir, List<DialogElement> dialogs)
{
if (string.IsNullOrEmpty(dialogDir) || dialogs == null) return false;
if (!File.Exists(dialogDir)) File.Create(dialogDir);
var texts = ParseDialogElements(dialogs);
File.WriteAllLines(dialogDir, texts);
return true;
}
private bool SaveTruncatedStates(string stateDir, List<StateKeyValue> states)
{
if (string.IsNullOrEmpty(stateDir) || states == null) return false;
if (!File.Exists(stateDir)) File.Create(stateDir);
var stateStr = JsonSerializer.Serialize(states, _options);
File.WriteAllText(stateDir, stateStr);
return true;
}
#endregion
}
}

View file

@ -114,14 +114,26 @@ public class ConversationController : ControllerBase
return response;
}
[HttpDelete("/conversation/{conversationId}/message/{messageId}")]
public async Task<bool> DeleteConversationMessage([FromRoute] string conversationId, [FromRoute] string messageId)
{
var conversationService = _services.GetRequiredService<IConversationService>();
var response = await conversationService.TruncateConversation(conversationId, messageId);
return response;
}
[HttpPost("/conversation/{agentId}/{conversationId}")]
public async Task<ChatResponseModel> SendMessage([FromRoute] string agentId,
[FromRoute] string conversationId,
[FromBody] NewMessageModel input)
{
var inputMsg = new RoleDialogModel(AgentRole.User, input.Text);
var conv = _services.GetRequiredService<IConversationService>();
if (!string.IsNullOrEmpty(input.TruncateMessageId))
{
await conv.TruncateConversation(conversationId, input.TruncateMessageId);
}
var inputMsg = new RoleDialogModel(AgentRole.User, input.Text);
conv.SetConversationId(conversationId, input.States);
conv.States.SetState("channel", input.Channel)
.SetState("provider", input.Provider)

View file

@ -263,4 +263,42 @@ public partial class MongoRepository
UpdatedTime = c.UpdatedTime
}).ToList();
}
public bool TruncateConversation(string conversationId, string messageId)
{
if (string.IsNullOrEmpty(conversationId) || string.IsNullOrEmpty(messageId)) return false;
var dialogFilter = Builders<ConversationDialogDocument>.Filter.Eq(x => x.ConversationId, conversationId);
var foundDialog = _dc.ConversationDialogs.Find(dialogFilter).FirstOrDefault();
if (foundDialog == null || foundDialog.Dialogs.IsNullOrEmpty()) return false;
var foundIdx = foundDialog.Dialogs.FindIndex(x => x.MetaData?.MessageId == messageId);
if (foundIdx < 0) return false;
// Handle truncated dialogs
var truncatedDialogs = foundDialog.Dialogs.Where((x, idx) => idx < foundIdx).ToList();
// Handle truncated states
var refTime = foundDialog.Dialogs.ElementAt(foundIdx).MetaData.CreateTime;
var stateFilter = Builders<ConversationStateDocument>.Filter.Eq(x => x.ConversationId, conversationId);
var foundStates = _dc.ConversationStates.Find(stateFilter).FirstOrDefault();
if (foundStates == null || foundStates.States.IsNullOrEmpty()) return false;
var truncatedStates = new List<StateMongoElement>();
foreach (var state in foundStates.States)
{
var values = state.Values.Where(x => x.UpdateTime < refTime).ToList();
if (values.Count == 0) continue;
state.Values = values;
truncatedStates.Add(state);
}
// Save
foundDialog.Dialogs = truncatedDialogs;
foundStates.States = truncatedStates;
_dc.ConversationDialogs.ReplaceOne(dialogFilter, foundDialog);
_dc.ConversationStates.ReplaceOne(stateFilter, foundStates);
return true;
}
}