refine state

This commit is contained in:
Jicheng Lu 2024-03-26 16:57:13 -05:00
parent b0eb2406a7
commit d572f90502
13 changed files with 105 additions and 79 deletions

View file

@ -1,6 +1,6 @@
namespace BotSharp.Abstraction.Conversations.Models;
public class ConversationState : Dictionary<string, List<StateValue>>
public class ConversationState : Dictionary<string, StateKeyValue>
{
public ConversationState()
{
@ -11,7 +11,7 @@ public class ConversationState : Dictionary<string, List<StateValue>>
{
foreach (var pair in pairs)
{
this[pair.Key] = pair.Values;
this[pair.Key] = pair;
}
}
}

View file

@ -3,6 +3,7 @@ namespace BotSharp.Abstraction.Conversations.Models;
public class StateKeyValue
{
public string Key { get; set; }
public bool Versioning { get; set; }
public List<StateValue> Values { get; set; } = new List<StateValue>();
public StateKeyValue()
@ -20,6 +21,10 @@ public class StateKeyValue
public class StateValue
{
public string Data { get; set; }
[JsonPropertyName("message_id")]
public string MessageId { get; set; }
public bool Active { get; set; }
[JsonPropertyName("update_time")]
public DateTime UpdateTime { get; set; }
public StateValue()

View file

@ -27,6 +27,7 @@ public partial class ConversationService
#endif
message.CurrentAgentId = agent.Id;
message.CreatedAt = DateTime.UtcNow;
if (string.IsNullOrEmpty(message.SenderId))
{
message.SenderId = _user.Id;

View file

@ -1,3 +1,5 @@
using BotSharp.Abstraction.Conversations.Models;
namespace BotSharp.Core.Conversations.Services;
/// <summary>
@ -42,12 +44,12 @@ public class ConversationStateService : IConversationStateService, IDisposable
var currentValue = value.ToString();
var hooks = _services.GetServices<IConversationHook>();
if (_states.TryGetValue(name, out var values))
if (ContainsState(name) && _states.TryGetValue(name, out var pair))
{
preValue = values?.LastOrDefault()?.Data ?? string.Empty;
preValue = pair?.Values.LastOrDefault()?.Data ?? string.Empty;
}
if (!_states.ContainsKey(name) || preValue != currentValue)
if (!ContainsState(name) || preValue != currentValue)
{
_logger.LogInformation($"[STATE] {name} = {value}");
foreach (var hook in hooks)
@ -55,19 +57,29 @@ public class ConversationStateService : IConversationStateService, IDisposable
hook.OnStateChanged(name, preValue, currentValue).Wait();
}
var stateValue = new StateValue
var routingCtx = _services.GetRequiredService<IRoutingContext>();
var newPair = new StateKeyValue
{
Data = currentValue,
UpdateTime = DateTime.UtcNow
Key = name,
Versioning = isNeedVersion
};
if (!_states.ContainsKey(name) || !isNeedVersion)
var newValue = new StateValue
{
_states[name] = new List<StateValue> { stateValue };
Data = currentValue,
MessageId = routingCtx.MessageId,
Active = true,
UpdateTime = DateTime.UtcNow,
};
if (!isNeedVersion || !_states.ContainsKey(name))
{
newPair.Values = new List<StateValue> { newValue };
_states[name] = newPair;
}
else
{
_states[name].Add(stateValue);
_states[name].Values.Add(newValue);
}
}
@ -85,9 +97,12 @@ public class ConversationStateService : IConversationStateService, IDisposable
{
foreach (var state in _states)
{
var value = state.Value?.LastOrDefault()?.Data ?? string.Empty;
curStates[state.Key] = value;
_logger.LogInformation($"[STATE] {state.Key} : {value}");
var value = state.Value?.Values?.LastOrDefault();
if (value == null || !value.Active) continue;
var data = value.Data ?? string.Empty;
curStates[state.Key] = data;
_logger.LogInformation($"[STATE] {state.Key} : {data}");
}
}
@ -112,7 +127,7 @@ public class ConversationStateService : IConversationStateService, IDisposable
foreach (var dic in _states)
{
states.Add(new StateKeyValue(dic.Key, dic.Value));
states.Add(dic.Value);
}
_db.UpdateConversationStates(_conversationId, states);
@ -121,7 +136,18 @@ public class ConversationStateService : IConversationStateService, IDisposable
public void CleanStates()
{
_states.Clear();
var utcNow = DateTime.UtcNow;
foreach (var key in _states.Keys)
{
var value = _states[key];
if (value == null || !value.Versioning || value.Values.IsNullOrEmpty()) continue;
var lastValue = value.Values.LastOrDefault();
if (lastValue == null || !lastValue.Active) continue;
lastValue.Active = false;
lastValue.UpdateTime = utcNow;
}
}
public Dictionary<string, string> GetStates()
@ -129,19 +155,22 @@ public class ConversationStateService : IConversationStateService, IDisposable
var curStates = new Dictionary<string, string>();
foreach (var state in _states)
{
curStates[state.Key] = state.Value?.LastOrDefault()?.Data ?? string.Empty;
var value = state.Value?.Values?.LastOrDefault();
if (value == null || !value.Active) continue;
curStates[state.Key] = value.Data ?? string.Empty;
}
return curStates;
}
public string GetState(string name, string defaultValue = "")
{
if (!_states.ContainsKey(name) || _states[name].IsNullOrEmpty())
if (!_states.ContainsKey(name) || _states[name].Values.IsNullOrEmpty() || !_states[name].Values.Last().Active)
{
return defaultValue;
}
return _states[name].Last().Data;
return _states[name].Values.Last().Data;
}
public void Dispose()
@ -152,8 +181,9 @@ public class ConversationStateService : IConversationStateService, IDisposable
public bool ContainsState(string name)
{
return _states.ContainsKey(name)
&& !_states[name].IsNullOrEmpty()
&& !string.IsNullOrEmpty(_states[name].Last().Data);
&& !_states[name].Values.IsNullOrEmpty()
&& _states[name].Values.LastOrDefault()?.Active == true
&& !string.IsNullOrEmpty(_states[name].Values.Last().Data);
}
public void SaveStateByArgs(JsonDocument args)

View file

@ -46,7 +46,7 @@ public class ConversationStorage : IConversationStorage
AgentId = agentId,
MessageId = dialog.MessageId,
FunctionName = dialog.FunctionName,
CreateTime = DateTime.UtcNow
CreateTime = dialog.CreatedAt
};
var content = dialog.Content.RemoveNewLine();
@ -65,7 +65,7 @@ public class ConversationStorage : IConversationStorage
MessageId = dialog.MessageId,
SenderId = dialog.SenderId,
FunctionName = dialog.FunctionName,
CreateTime = DateTime.UtcNow
CreateTime = dialog.CreatedAt
};
var content = dialog.Content.RemoveNewLine();

View file

@ -481,7 +481,7 @@ namespace BotSharp.Core.Repository
var refTime = dialogs.ElementAt(foundIdx).MetaData.CreateTime;
var stateDir = Path.Combine(convDir, STATE_FILE);
var states = CollectConversationStates(stateDir);
isSaved = HandleTruncatedStates(stateDir, states, refTime);
isSaved = HandleTruncatedStates(stateDir, states, messageId, refTime);
// Handle truncated breakpoints
var breakpointDir = Path.Combine(convDir, BREAKPOINT_FILE);
@ -597,12 +597,20 @@ namespace BotSharp.Core.Repository
return isSaved;
}
private bool HandleTruncatedStates(string stateDir, List<StateKeyValue> states, DateTime refTime)
private bool HandleTruncatedStates(string stateDir, List<StateKeyValue> states, string refMsgId, DateTime refTime)
{
var truncatedStates = new List<StateKeyValue>();
foreach (var state in states)
{
var values = state.Values.Where(x => x.UpdateTime < refTime).ToList();
if (!state.Versioning)
{
truncatedStates.Add(state);
continue;
}
var values = state.Values.Where(x => x.MessageId != refMsgId)
.Where(x => x.UpdateTime < refTime)
.ToList();
if (values.Count == 0) continue;
state.Values = values;

View file

@ -71,11 +71,7 @@ namespace BotSharp.Core.Repository
log.MessageId = log.MessageId.IfNullOrEmptyAs(Guid.NewGuid().ToString());
var convDir = FindConversationDirectory(log.ConversationId);
if (string.IsNullOrEmpty(convDir))
{
convDir = Path.Combine(_dbSettings.FileRepository, _conversationSettings.DataDir, log.ConversationId);
Directory.CreateDirectory(convDir);
}
if (string.IsNullOrEmpty(convDir)) return;
var logDir = Path.Combine(convDir, "content_log");
if (!Directory.Exists(logDir))
@ -120,11 +116,7 @@ namespace BotSharp.Core.Repository
log.MessageId = log.MessageId.IfNullOrEmptyAs(Guid.NewGuid().ToString());
var convDir = FindConversationDirectory(log.ConversationId);
if (string.IsNullOrEmpty(convDir))
{
convDir = Path.Combine(_dbSettings.FileRepository, _conversationSettings.DataDir, log.ConversationId);
Directory.CreateDirectory(convDir);
}
if (string.IsNullOrEmpty(convDir)) return;
var logDir = Path.Combine(convDir, "state_log");
if (!Directory.Exists(logDir))

View file

@ -36,7 +36,6 @@ namespace BotSharp.OpenAPI.BackgroundServices
try
{
await CleanIdleConversationsAsync();
await CloseIdleConversationsAsync(TimeSpan.FromMinutes(10));
}
catch (Exception ex)
{
@ -54,37 +53,6 @@ namespace BotSharp.OpenAPI.BackgroundServices
await base.StopAsync(stoppingToken);
}
private async Task CloseIdleConversationsAsync(TimeSpan conversationIdleTimeout)
{
using var scope = _services.CreateScope();
var conversationService = scope.ServiceProvider.GetRequiredService<IConversationService>();
var hooks = scope.ServiceProvider.GetServices<IConversationHook>()
.OrderBy(x => x.Priority)
.ToList();
var moment = DateTime.UtcNow.Add(-conversationIdleTimeout);
var conversations = (await conversationService.GetLastConversations()).Where(c => c.CreatedTime <= moment);
foreach (var conversation in conversations)
{
try
{
var response = new RoleDialogModel(AgentRole.Assistant, "End the conversation due to timeout.")
{
StopCompletion = true,
FunctionName = "conversation_end"
};
foreach (var hook in hooks)
{
await hook.OnConversationEnding(response);
}
}
catch (Exception ex)
{
_logger.LogError(ex, $"Error occurred closing conversation #{conversation.Id}.");
}
}
}
private async Task CleanIdleConversationsAsync()
{
using var scope = _services.CreateScope();

View file

@ -1,3 +1,4 @@
using BotSharp.Abstraction.Routing;
using BotSharp.Abstraction.Users.Models;
namespace BotSharp.OpenAPI.Controllers;
@ -162,6 +163,10 @@ public class ConversationController : ControllerBase
var inputMsg = new RoleDialogModel(AgentRole.User, input.Text);
conv.SetConversationId(conversationId, input.States);
var routing = _services.GetRequiredService<IRoutingService>();
routing.Context.SetMessageId(conversationId, inputMsg.MessageId);
conv.States.SetState("channel", input.Channel)
.SetState("provider", input.Provider)
.SetState("model", input.Model)

View file

@ -345,7 +345,7 @@ public class StreamingLogHook : ConversationHookBase, IContentGeneratingHook, IR
Role = input.Message.Role,
Content = input.Log,
Source = input.Source,
CreateTime = DateTime.UtcNow
CreateTime = input.Message.CreatedAt
};
var json = JsonSerializer.Serialize(output, _options.JsonSerializerOptions);
@ -367,7 +367,7 @@ public class StreamingLogHook : ConversationHookBase, IContentGeneratingHook, IR
ConversationId = conversationId,
MessageId = message.MessageId,
States = states,
CreateTime = DateTime.UtcNow
CreateTime = message.CreatedAt
};
var convSettings = _services.GetRequiredService<ConversationSetting>();

View file

@ -5,6 +5,7 @@ namespace BotSharp.Plugin.MongoStorage.Models;
public class StateMongoElement
{
public string Key { get; set; }
public bool Versioning { get; set; }
public List<StateValueMongoElement> Values { get; set; }
public static StateMongoElement ToMongoElement(StateKeyValue state)
@ -12,6 +13,7 @@ public class StateMongoElement
return new StateMongoElement
{
Key = state.Key,
Versioning = state.Versioning,
Values = state.Values?.Select(x => StateValueMongoElement.ToMongoElement(x))?.ToList() ?? new List<StateValueMongoElement>()
};
}
@ -21,6 +23,7 @@ public class StateMongoElement
return new StateKeyValue
{
Key = state.Key,
Versioning = state.Versioning,
Values = state.Values?.Select(x => StateValueMongoElement.ToDomainElement(x))?.ToList() ?? new List<StateValue>()
};
}
@ -29,6 +32,8 @@ public class StateMongoElement
public class StateValueMongoElement
{
public string Data { get; set; }
public string MessageId { get; set; }
public bool Active { get; set; }
public DateTime UpdateTime { get; set; }
public static StateValueMongoElement ToMongoElement(StateValue element)
@ -36,6 +41,8 @@ public class StateValueMongoElement
return new StateValueMongoElement
{
Data = element.Data,
MessageId = element.MessageId,
Active = element.Active,
UpdateTime = element.UpdateTime
};
}
@ -45,6 +52,8 @@ public class StateValueMongoElement
return new StateValue
{
Data = element.Data,
MessageId = element.MessageId,
Active = element.Active,
UpdateTime = element.UpdateTime
};
}

View file

@ -453,7 +453,15 @@ public partial class MongoRepository
var truncatedStates = new List<StateMongoElement>();
foreach (var state in foundStates.States)
{
var values = state.Values.Where(x => x.UpdateTime < refTime).ToList();
if (!state.Versioning)
{
truncatedStates.Add(state);
continue;
}
var values = state.Values.Where(x => x.MessageId != messageId)
.Where(x => x.UpdateTime < refTime)
.ToList();
if (values.Count == 0) continue;
state.Values = values;

View file

@ -64,13 +64,13 @@ public partial class MongoRepository
{
if (log == null) return;
var conversationId = log.ConversationId.IfNullOrEmptyAs(Guid.NewGuid().ToString());
var messageId = log.MessageId.IfNullOrEmptyAs(Guid.NewGuid().ToString());
var found = _dc.Conversations.AsQueryable().FirstOrDefault(x => x.Id == log.ConversationId);
if (found == null) return;
var logDoc = new ConversationContentLogDocument
{
ConversationId = conversationId,
MessageId = messageId,
ConversationId = log.ConversationId,
MessageId = log.MessageId,
Name = log.Name,
AgentId = log.AgentId,
Role = log.Role,
@ -109,13 +109,13 @@ public partial class MongoRepository
{
if (log == null) return;
var conversationId = log.ConversationId.IfNullOrEmptyAs(Guid.NewGuid().ToString());
var messageId = log.MessageId.IfNullOrEmptyAs(Guid.NewGuid().ToString());
var found = _dc.Conversations.AsQueryable().FirstOrDefault(x => x.Id == log.ConversationId);
if (found == null) return;
var logDoc = new ConversationStateLogDocument
{
ConversationId = conversationId,
MessageId = messageId,
ConversationId = log.ConversationId,
MessageId = log.MessageId,
States = log.States,
CreateTime = log.CreateTime
};