diff --git a/Bot.WebStarter/settings.bot.json b/Bot.WebStarter/settings.bot.json index 96cf8614..30f3bb06 100644 --- a/Bot.WebStarter/settings.bot.json +++ b/Bot.WebStarter/settings.bot.json @@ -1,5 +1,5 @@ { "Rasa": { - "Host": "http://rasa.local:5000" + "Host": "http://gtx.local:5000" } } diff --git a/BotSharp.Core/Agents/AgentExtension.cs b/BotSharp.Core/Agents/AgentExtension.cs index 42c44602..cceebb39 100644 --- a/BotSharp.Core/Agents/AgentExtension.cs +++ b/BotSharp.Core/Agents/AgentExtension.cs @@ -30,7 +30,78 @@ namespace BotSharp.Core.Agents return entity.Id; } - public static RasaTrainingData GrabCorpus(this Agent agent, Database dc, List ctx) + public static RasaTrainingData GrabCorpus(this Agent agent, Database dc) + { + var trainingData = new RasaTrainingData + { + Entities = new List(), + UserSays = new List() + }; + + var expressParts = new List(); + + var intents = dc.Table() + .Include(x => x.Contexts) + .Include(x => x.UserSays).ThenInclude(say => say.Data) + .Where(x => x.UserSays.Count > 0) + .ToList(); + + intents.ForEach(intent => + { + intent.UserSays.ForEach(exp => + { + var say = new RasaIntentExpression + { + Intent = intent.Name, + Text = String.Join("", exp.Data.OrderBy(x => x.UpdatedTime).Select(x => x.Text)), + }; + + // convert entity format + exp.Data.Where(x => !String.IsNullOrEmpty(x.Meta)) + .ToList() + .ForEach(x => + { + int start = say.Text.IndexOf(x.Text); + + var part = new RasaIntentExpressionPart + { + Value = x.Text, + Entity = x.Alias, + Start = start, + End = start + x.Text.Length + }; + + if (say.Entities == null) say.Entities = new List(); + say.Entities.Add(part); + + // assemble entity synonmus + if (!trainingData.Entities.Any(y => y.EntityType == x.Alias && y.EntityValue == x.Text)) + { + var allSynonyms = (from e in dc.Table() + join ee in dc.Table() on e.Id equals ee.EntityId + join ees in dc.Table() on ee.Id equals ees.EntityEntryId + where e.Name == x.Alias && ee.Value == x.Text & ees.Synonym != x.Text + select ees.Synonym).ToList(); + + var te = new RasaTraningEntity + { + EntityType = x.Alias, + EntityValue = x.Text, + Synonyms = allSynonyms + }; + + trainingData.Entities.Add(te); + } + }); + + trainingData.UserSays.Add(say); + }); + }); + + return trainingData; + } + + public static RasaTrainingData GrabCorpusPerContexts(this Agent agent, Database dc, List ctx) { var trainingData = new RasaTrainingData { diff --git a/BotSharp.Core/Engines/RequestExtension.cs b/BotSharp.Core/Engines/RequestExtension.cs index a6deb864..2b25f9f2 100644 --- a/BotSharp.Core/Engines/RequestExtension.cs +++ b/BotSharp.Core/Engines/RequestExtension.cs @@ -54,61 +54,11 @@ namespace BotSharp.Core.Engines } }).OrderByDescending(x => x.Contexts.Count).ToList(); - // training per request contexts - { - string contextId = $"{String.Join(',', contexts.Select(x => x.Name))}".GetMd5Hash(); - string modelName = dc.Table().FirstOrDefault(x => x.ContextId == contextId)?.ModelName; - // need training - if (String.IsNullOrEmpty(modelName)) - { - request.Contexts = contexts.Select(x => new AIContext { Name = x.Name.ToLower() }) - .OrderBy(x => x.Name) - .ToList(); + var result = CallRasa(rasa.agent.Id, request.Query.First(), rasa.agent.Id); - dc.DbTran(() => - { - modelName = TrainWithContexts(rasa, dc, request, contextId); - }); - } - - var result = CallRasa(rasa.agent.Id, request.Query.First(), modelName); - - if (result.Data.Intent != null) - { - response = result.Data; - } - } - - // Max contexts match - if (response == null) - { - foreach (var it in intents) - { - request.Contexts = it.Contexts.Select(x => new AIContext { Name = x.Name.ToLower() }) - .OrderBy(x => x.Name) - .ToList(); - string contextId = $"{String.Join(',', request.Contexts.Select(x => x.Name))}".GetMd5Hash(); - - string modelName = dc.Table().FirstOrDefault(x => x.ContextId == contextId)?.ModelName; - - // need training - if (String.IsNullOrEmpty(modelName)) - { - dc.DbTran(() => - { - modelName = TrainWithContexts(rasa, dc, request, contextId); - }); - } - - var result = CallRasa(rasa.agent.Id, request.Query.First(), modelName); - - if (result.Data.Intent != null) - { - response = result.Data; - break; - } - }; - } + result.Data.IntentRanking = result.Data.IntentRanking.Where(x => intents.Select(i => i.Name).Contains(x.Name)).ToList(); + result.Data.Intent = result.Data.IntentRanking.First(); + response = result.Data; var intent = (dc.Table().Where(x => x.Name == response.Intent.Name) .Include(x => x.Responses).ThenInclude(x => x.Contexts) @@ -121,6 +71,28 @@ namespace BotSharp.Core.Engines aiResponse.Status = new AIResponseStatus { }; aiResponse.SessionId = rasa.AiConfig.SessionId; aiResponse.Timestamp = DateTime.UtcNow; + intentResponse.Parameters.ForEach(p => { + string query = request.Query.First(); + var entity = response.Entities.FirstOrDefault(x => x.Entity == p.Name); + if(entity != null) + { + p.Value = query.Substring(entity.Start, entity.End - entity.Start); + } + + // fixed entity per request + if(request.Entities != null) + { + var fixedEntity = request.Entities.FirstOrDefault(x => x.Name == p.Name); + if (fixedEntity != null) + { + if (query.ToLower().Contains(fixedEntity.Entries.First().Value.ToLower())) + { + p.Value = fixedEntity.Entries.First().Value; + } + } + } + + }); intentResponse.Messages = intentResponse.Messages.OrderBy(x => x.UpdatedTime).ToList(); intentResponse.Messages.ToList() .ForEach(msg => @@ -142,7 +114,7 @@ namespace BotSharp.Core.Engines Source = "agent", ResolvedQuery = request.Query.First(), Action = intentResponse.Action, - Parameters = new Dictionary(), + Parameters = intentResponse.Parameters.ToDictionary(x => x.Name, x=> x.Value), Score = response.Intent.Confidence, Metadata = new AIResponseMetadata { IntentId = intent.Id, IntentName = intent.Name }, Fulfillment = new AIResponseFulfillment @@ -225,6 +197,252 @@ namespace BotSharp.Core.Engines return client.Execute(rest); } + + public static AIResponse TextRequestPerContexts(this RasaAi rasa, AIRequest request) + { + AIResponse aiResponse = new AIResponse(); + RasaResponse response = null; + Database dc = rasa.dc; + + // Merge input contexts + var contexts = dc.Table() + .Where(x => x.SessionId == rasa.AiConfig.SessionId && x.Lifespan > 0) + .ToList() + .Select(x => new AIContext { Name = x.Context.ToLower(), Lifespan = x.Lifespan }) + .ToList(); + + contexts.AddRange(request.Contexts.Select(x => new AIContext { Name = x.Name.ToLower(), Lifespan = x.Lifespan })); + contexts = contexts.OrderBy(x => x.Name).ToList(); + + // search all potential intents which input context included in contexts + var intents = rasa.agent.Intents.Where(it => + { + if (contexts.Count == 0) + { + return it.Contexts.Count() == 0; + } + else + { + return it.Contexts.Count() > 0 && + it.Contexts.Count(x => contexts.Select(ctx => ctx.Name).Contains(x.Name.ToLower())) == it.Contexts.Count; + } + }).OrderByDescending(x => x.Contexts.Count).ToList(); + + // training per request contexts + { + string contextId = $"{String.Join(',', contexts.Select(x => x.Name))}".GetMd5Hash(); + string modelName = dc.Table().FirstOrDefault(x => x.ContextId == contextId)?.ModelName; + // need training + if (String.IsNullOrEmpty(modelName)) + { + request.Contexts = contexts.Select(x => new AIContext { Name = x.Name.ToLower() }) + .OrderBy(x => x.Name) + .ToList(); + + dc.DbTran(() => + { + modelName = TrainWithContexts(rasa, dc, request, contextId); + }); + } + + var result = CallRasa(rasa.agent.Id, request.Query.First(), modelName); + + if (result.Data.Intent != null) + { + response = result.Data; + } + } + + // Max contexts match + if (response == null) + { + foreach (var it in intents) + { + request.Contexts = it.Contexts.Select(x => new AIContext { Name = x.Name.ToLower() }) + .OrderBy(x => x.Name) + .ToList(); + string contextId = $"{String.Join(',', request.Contexts.Select(x => x.Name))}".GetMd5Hash(); + + string modelName = dc.Table().FirstOrDefault(x => x.ContextId == contextId)?.ModelName; + + // need training + if (String.IsNullOrEmpty(modelName)) + { + dc.DbTran(() => + { + modelName = TrainWithContexts(rasa, dc, request, contextId); + }); + } + + var result = CallRasa(rasa.agent.Id, request.Query.First(), modelName); + + if (result.Data.Intent != null) + { + response = result.Data; + break; + } + }; + } + + var intent = (dc.Table().Where(x => x.Name == response.Intent.Name) + .Include(x => x.Responses).ThenInclude(x => x.Contexts) + .Include(x => x.Responses).ThenInclude(x => x.Parameters) + .Include(x => x.Responses).ThenInclude(x => x.Messages)).First(); + + var intentResponse = ArrayHelper.GetRandom(intent.Responses); + aiResponse.Id = Guid.NewGuid().ToString(); + aiResponse.Lang = rasa.agent.Language; + aiResponse.Status = new AIResponseStatus { }; + aiResponse.SessionId = rasa.AiConfig.SessionId; + aiResponse.Timestamp = DateTime.UtcNow; + intentResponse.Messages = intentResponse.Messages.OrderBy(x => x.UpdatedTime).ToList(); + intentResponse.Messages.ToList() + .ForEach(msg => + { + if (msg.Type == AIResponseMessageType.Custom) + { + + } + else + { + msg.Speech = msg.Speech.StartsWith("[") ? + ArrayHelper.GetRandom(msg.Speech.Substring(2, msg.Speech.Length - 4).Split("\",\"").ToList()) : + msg.Speech; + } + }); + + aiResponse.Result = new AIResponseResult + { + Source = "agent", + ResolvedQuery = request.Query.First(), + Action = intentResponse.Action, + Parameters = new Dictionary(), + Score = response.Intent.Confidence, + Metadata = new AIResponseMetadata { IntentId = intent.Id, IntentName = intent.Name }, + Fulfillment = new AIResponseFulfillment + { + Messages = intentResponse.Messages.Select(x => { + if (x.Type == AIResponseMessageType.Custom) + { + return (new + { + x.Type, + Payload = JObject.Parse(x.Payload) + }) as Object; + } + else + { + return (new { x.Type, x.Speech }) as Object; + } + + }).ToList() + } + }; + + // Merge context lifespan + // override if exists, otherwise add, delete if lifespan is zero + dc.DbTran(() => + { + var sessionContexts = dc.Table().Where(x => x.SessionId == rasa.AiConfig.SessionId).ToList(); + + // minus 1 round + sessionContexts.Where(x => !intentResponse.Contexts.Select(ctx => ctx.Name).Contains(x.Context)) + .ToList() + .ForEach(ctx => ctx.Lifespan = ctx.Lifespan - 1); + + intentResponse.Contexts.ForEach(ctx => + { + var session1 = sessionContexts.FirstOrDefault(x => x.Context == ctx.Name); + + if (session1 != null) + { + if (ctx.Lifespan == 0) + { + dc.Table().Remove(session1); + } + else + { + session1.Lifespan = ctx.Lifespan; + } + } + else + { + dc.Table().Add(new SessionContext + { + SessionId = rasa.AiConfig.SessionId, + Context = ctx.Name, + Lifespan = ctx.Lifespan + }); + } + }); + }); + + aiResponse.Result.Contexts = dc.Table() + .Where(x => x.SessionId == rasa.AiConfig.SessionId) + .Select(x => new AIContext { Name = x.Context.ToLower(), Lifespan = x.Lifespan }) + .ToArray(); + + return aiResponse; + } + + public static string Train(this RasaAi console, Database dc) + { + var corpus = console.agent.GrabCorpus(dc); + + // Add some fake data + if(corpus.UserSays.Count < 3) + { + corpus.UserSays.Add(new RasaIntentExpression + { + Intent = "Welcome", + Text = "Hi" + }); + + corpus.UserSays.Add(new RasaIntentExpression + { + Intent = "Welcome", + Text = "Hey" + }); + + corpus.UserSays.Add(new RasaIntentExpression + { + Intent = "Welcome", + Text = "Hello" + }); + } + + string json = JsonConvert.SerializeObject(new { rasa_nlu_data = corpus }, + new JsonSerializerSettings + { + ContractResolver = new CamelCasePropertyNamesContractResolver(), + NullValueHandling = NullValueHandling.Ignore + }); + + var client = new RestClient($"{Database.Configuration.GetSection("Rasa:Host").Value}"); + var rest = new RestRequest("train", Method.POST); + rest.AddQueryParameter("project", console.agent.Id); + rest.AddParameter("application/json", json, ParameterType.RequestBody); + + var response = client.Execute(rest); + + if (response.IsSuccessful) + { + var result = JObject.Parse(response.Content); + + string modelName = result["info"].Value().Split(": ")[1]; + + return modelName; + } + else + { + var result = JObject.Parse(response.Content); + + Console.WriteLine(result["error"]); + + return String.Empty; + } + } + /// /// Need two categories at least /// @@ -235,7 +453,7 @@ namespace BotSharp.Core.Engines /// public static string TrainWithContexts(this RasaAi console, Database dc, AIRequest request, String contextId) { - var corpus = console.agent.GrabCorpus(dc, request.Contexts); + var corpus = console.agent.GrabCorpusPerContexts(dc, request.Contexts); corpus.UserSays.Add(new RasaIntentExpression { diff --git a/BotSharp.Core/Models/AIResponseResult.cs b/BotSharp.Core/Models/AIResponseResult.cs index 306c5735..4fc59506 100644 --- a/BotSharp.Core/Models/AIResponseResult.cs +++ b/BotSharp.Core/Models/AIResponseResult.cs @@ -31,7 +31,7 @@ namespace BotSharp.Core.Models } } - public Dictionary Parameters { get; set; } + public Dictionary Parameters { get; set; } public AIContext[] Contexts { get; set; } @@ -125,10 +125,10 @@ namespace BotSharp.Core.Models if (Parameters.ContainsKey(name)) { - var parameter = Parameters[name] as JObject; + var parameter = Parameters[name].ToString(); if (parameter != null) { - return parameter; + return JObject.FromObject(parameter); } } diff --git a/BotSharp.UnitTest/AgentTest.cs b/BotSharp.UnitTest/AgentTest.cs index a1a59293..52097c1b 100644 --- a/BotSharp.UnitTest/AgentTest.cs +++ b/BotSharp.UnitTest/AgentTest.cs @@ -14,7 +14,7 @@ namespace BotSharp.UnitTest public class AgentTest : TestEssential { [TestMethod] - public void CreateAgent() + public void CreateAgentTes() { var agent = new Agent { @@ -28,7 +28,7 @@ namespace BotSharp.UnitTest } [TestMethod] - public void UpdateAgent() + public void UpdateAgentTest() { var agent = new Agent { @@ -42,7 +42,7 @@ namespace BotSharp.UnitTest } [TestMethod] - public void RestoreAgent() + public void RestoreAgentTest() { var rasa = new RasaAi(dc); var importer = new AgentImporterInDialogflow(); @@ -56,9 +56,14 @@ namespace BotSharp.UnitTest } [TestMethod] - public void Train() + public void TrainAgentTest() { - var rasa = new RasaAi(dc); + var config = new AIConfiguration(BOT_CLIENT_TOKEN, SupportedLanguage.English); + config.SessionId = Guid.NewGuid().ToString(); + + var rasa = new RasaAi(dc, config); + rasa.agent = rasa.LoadAgent(); + rasa.Train(dc); } } } diff --git a/BotSharp.UnitTest/BotSharp.UnitTest.csproj b/BotSharp.UnitTest/BotSharp.UnitTest.csproj index c0466d62..5d65ea16 100644 --- a/BotSharp.UnitTest/BotSharp.UnitTest.csproj +++ b/BotSharp.UnitTest/BotSharp.UnitTest.csproj @@ -9,9 +9,9 @@ - - - + + + diff --git a/BotSharp.UnitTest/IntentTest.cs b/BotSharp.UnitTest/IntentTest.cs index 4b1bfad1..75d9864b 100644 --- a/BotSharp.UnitTest/IntentTest.cs +++ b/BotSharp.UnitTest/IntentTest.cs @@ -18,11 +18,8 @@ namespace BotSharp.UnitTest var rasa = new RasaAi(dc, config); - var response = rasa.TextRequest(new AIRequest { Query = new String[] { "Create a work order for PetSmart" } }); - Assert.IsTrue(response.Result.Metadata.IntentName == "Create Work Order"); - - response = rasa.TextRequest(new AIRequest { Query = new String[] { "1010" } }); - Assert.IsTrue(response.Result.Metadata.IntentName == "Telling Store Number"); + var response = rasa.TextRequest(new AIRequest { Query = new String[] { "Hello" } }); + Assert.IsTrue(response.Result.Metadata.IntentName == "greeting"); } } }