refine state
This commit is contained in:
parent
b0eb2406a7
commit
d572f90502
|
|
@ -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;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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>();
|
||||
|
|
|
|||
|
|
@ -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
|
||||
};
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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
|
||||
};
|
||||
|
|
|
|||
Loading…
Reference in a new issue