BotSharp/BotSharp.Core/Engines/RequestExtension.cs

516 lines
20 KiB
C#
Raw Normal View History

2018-03-28 22:08:49 +00:00
using BotSharp.Core.Agents;
using BotSharp.Core.Intents;
using BotSharp.Core.Models;
using BotSharp.Core.Sessions;
using DotNetToolkit;
2017-12-30 20:26:35 +00:00
using EntityFrameworkCore.BootKit;
using Microsoft.EntityFrameworkCore;
2017-12-18 05:30:20 +00:00
using Newtonsoft.Json;
using Newtonsoft.Json.Linq;
2017-12-18 05:30:20 +00:00
using Newtonsoft.Json.Serialization;
using RestSharp;
using System;
using System.Collections.Generic;
using System.IO;
using System.Linq;
2017-12-18 05:30:20 +00:00
using System.Text;
2018-03-28 22:08:49 +00:00
namespace BotSharp.Core.Engines
2017-12-18 05:30:20 +00:00
{
public static class RequestExtension
{
2018-03-28 22:08:49 +00:00
public static AIResponse TextRequest(this RasaAi rasa, string text, RequestExtras requestExtras)
{
return rasa.TextRequest(new AIRequest(text, requestExtras));
}
public static AIResponse TextRequest(this RasaAi rasa, AIRequest request)
2017-12-18 05:30:20 +00:00
{
AIResponse aiResponse = new AIResponse();
RasaResponse response = null;
2018-03-28 22:08:49 +00:00
Database dc = rasa.dc;
2017-12-18 05:30:20 +00:00
// Merge input contexts
var contexts = dc.Table<SessionContext>()
2018-03-28 22:08:49 +00:00
.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)
2017-12-30 20:26:35 +00:00
{
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();
2018-04-06 22:39:30 +00:00
var result = CallRasa(rasa.agent.Id, request.Query.First(), rasa.agent.Id);
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<Intent>().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.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 =>
{
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 = 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
{
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<SessionContext>().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<SessionContext>().Remove(session1);
}
else
{
session1.Lifespan = ctx.Lifespan;
}
}
else
{
dc.Table<SessionContext>().Add(new SessionContext
{
SessionId = rasa.AiConfig.SessionId,
Context = ctx.Name,
Lifespan = ctx.Lifespan
});
}
});
});
aiResponse.Result.Contexts = dc.Table<SessionContext>()
.Where(x => x.SessionId == rasa.AiConfig.SessionId)
.Select(x => new AIContext { Name = x.Context.ToLower(), Lifespan = x.Lifespan })
.ToArray();
return aiResponse;
}
private static IRestResponse<RasaResponse> CallRasa(string projectId, string text, string model)
{
var client = new RestClient($"{Database.Configuration.GetSection("Rasa:Host").Value}");
var rest = new RestRequest("parse", Method.POST);
string json = JsonConvert.SerializeObject(new { Project = projectId, Q = text, Model = model },
new JsonSerializerSettings
{
ContractResolver = new CamelCasePropertyNamesContractResolver()
});
rest.AddParameter("application/json", json, ParameterType.RequestBody);
return client.Execute<RasaResponse>(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<SessionContext>()
.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();
2018-03-28 22:08:49 +00:00
// training per request contexts
{
2018-03-28 22:08:49 +00:00
string contextId = $"{String.Join(',', contexts.Select(x => x.Name))}".GetMd5Hash();
string modelName = dc.Table<ContextModelMapping>().FirstOrDefault(x => x.ContextId == contextId)?.ModelName;
// need training
if (String.IsNullOrEmpty(modelName))
{
2018-03-28 22:08:49 +00:00
request.Contexts = contexts.Select(x => new AIContext { Name = x.Name.ToLower() })
.OrderBy(x => x.Name)
.ToList();
dc.DbTran(() =>
{
modelName = TrainWithContexts(rasa, dc, request, contextId);
});
}
2018-03-28 22:08:49 +00:00
var result = CallRasa(rasa.agent.Id, request.Query.First(), modelName);
2018-03-28 22:08:49 +00:00
if (result.Data.Intent != null)
{
response = result.Data;
}
2018-03-28 22:08:49 +00:00
}
// 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<ContextModelMapping>().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<Intent>().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 { };
2018-03-28 22:08:49 +00:00
aiResponse.SessionId = rasa.AiConfig.SessionId;
aiResponse.Timestamp = DateTime.UtcNow;
2018-03-28 22:08:49 +00:00
intentResponse.Messages = intentResponse.Messages.OrderBy(x => x.UpdatedTime).ToList();
intentResponse.Messages.ToList()
.ForEach(msg =>
{
2018-03-28 22:08:49 +00:00
if (msg.Type == AIResponseMessageType.Custom)
{
2018-04-06 22:39:30 +00:00
2018-03-28 22:08:49 +00:00
}
else
{
msg.Speech = msg.Speech.StartsWith("[") ?
ArrayHelper.GetRandom(msg.Speech.Substring(2, msg.Speech.Length - 4).Split("\",\"").ToList()) :
msg.Speech;
}
2017-12-30 20:26:35 +00:00
});
2017-12-18 05:30:20 +00:00
aiResponse.Result = new AIResponseResult
{
Source = "agent",
ResolvedQuery = request.Query.First(),
Action = intentResponse.Action,
2018-04-06 22:39:30 +00:00
Parameters = new Dictionary<string, string>(),
Score = response.Intent.Confidence,
Metadata = new AIResponseMetadata { IntentId = intent.Id, IntentName = intent.Name },
Fulfillment = new AIResponseFulfillment
{
2018-03-28 22:08:49 +00:00
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;
}
2018-04-06 22:39:30 +00:00
2018-03-28 22:08:49 +00:00
}).ToList()
}
};
// Merge context lifespan
// override if exists, otherwise add, delete if lifespan is zero
dc.DbTran(() =>
{
2018-03-28 22:08:49 +00:00
var sessionContexts = dc.Table<SessionContext>().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<SessionContext>().Remove(session1);
}
else
{
session1.Lifespan = ctx.Lifespan;
}
}
else
{
dc.Table<SessionContext>().Add(new SessionContext
{
2018-03-28 22:08:49 +00:00
SessionId = rasa.AiConfig.SessionId,
Context = ctx.Name,
Lifespan = ctx.Lifespan
});
}
});
});
aiResponse.Result.Contexts = dc.Table<SessionContext>()
2018-03-28 22:08:49 +00:00
.Where(x => x.SessionId == rasa.AiConfig.SessionId)
.Select(x => new AIContext { Name = x.Context.ToLower(), Lifespan = x.Lifespan })
.ToArray();
2017-12-18 05:30:20 +00:00
return aiResponse;
2017-12-18 05:30:20 +00:00
}
2018-04-06 22:39:30 +00:00
public static string Train(this RasaAi console, Database dc)
2018-03-28 22:08:49 +00:00
{
2018-04-06 22:39:30 +00:00
var corpus = console.agent.GrabCorpus(dc);
2018-03-28 22:08:49 +00:00
2018-04-06 22:39:30 +00:00
// 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 },
2018-03-28 22:08:49 +00:00
new JsonSerializerSettings
{
2018-04-06 22:39:30 +00:00
ContractResolver = new CamelCasePropertyNamesContractResolver(),
NullValueHandling = NullValueHandling.Ignore
2018-03-28 22:08:49 +00:00
});
2018-04-06 22:39:30 +00:00
var client = new RestClient($"{Database.Configuration.GetSection("Rasa:Host").Value}");
var rest = new RestRequest("train", Method.POST);
rest.AddQueryParameter("project", console.agent.Id);
2018-03-28 22:08:49 +00:00
rest.AddParameter("application/json", json, ParameterType.RequestBody);
2018-04-06 22:39:30 +00:00
var response = client.Execute(rest);
if (response.IsSuccessful)
{
var result = JObject.Parse(response.Content);
string modelName = result["info"].Value<String>().Split(": ")[1];
return modelName;
}
else
{
var result = JObject.Parse(response.Content);
Console.WriteLine(result["error"]);
return String.Empty;
}
2018-03-28 22:08:49 +00:00
}
2018-04-06 22:39:30 +00:00
2017-12-30 20:26:35 +00:00
/// <summary>
/// Need two categories at least
/// </summary>
/// <param name="console"></param>
/// <param name="dc"></param>
/// <param name="request"></param>
/// <param name="contextId"></param>
2017-12-30 20:26:35 +00:00
/// <returns></returns>
public static string TrainWithContexts(this RasaAi console, Database dc, AIRequest request, String contextId)
2017-12-18 05:30:20 +00:00
{
2018-04-06 22:39:30 +00:00
var corpus = console.agent.GrabCorpusPerContexts(dc, request.Contexts);
2018-03-28 22:08:49 +00:00
corpus.UserSays.Add(new RasaIntentExpression
{
Intent = "Welcome",
Text = "Hi"
});
2018-03-28 22:08:49 +00:00
corpus.UserSays.Add(new RasaIntentExpression
{
Intent = "Welcome",
Text = "Hey"
});
2018-03-28 22:08:49 +00:00
corpus.UserSays.Add(new RasaIntentExpression
{
Intent = "Welcome",
Text = "Hello"
});
2017-12-18 05:30:20 +00:00
2017-12-18 13:31:15 +00:00
string json = JsonConvert.SerializeObject(new { rasa_nlu_data = corpus },
new JsonSerializerSettings
2017-12-18 05:30:20 +00:00
{
2018-04-02 22:42:11 +00:00
ContractResolver = new CamelCasePropertyNamesContractResolver(),
NullValueHandling = NullValueHandling.Ignore
2017-12-18 13:31:15 +00:00
});
2017-12-18 05:30:20 +00:00
2018-03-28 22:08:49 +00:00
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)
{
2018-04-02 22:42:11 +00:00
var result = JObject.Parse(response.Content);
string modelName = result["info"].Value<String>().Split(": ")[1];
dc.Table<ContextModelMapping>().Add(new ContextModelMapping
{
AgentId = console.agent.Id,
ModelName = modelName,
ContextId = contextId
});
2017-12-18 05:30:20 +00:00
return modelName;
}
else
{
2018-04-02 22:42:11 +00:00
var result = JObject.Parse(response.Content);
Console.WriteLine(result["error"]);
2017-12-18 05:30:20 +00:00
return String.Empty;
}
2017-12-18 05:30:20 +00:00
}
}
}