BotSharp/src/Infrastructure/BotSharp.Core/Translation/TranslationService.cs

356 lines
11 KiB
C#
Raw Normal View History

2024-05-09 20:59:59 +00:00
using BotSharp.Abstraction.Infrastructures.Enums;
2024-04-25 16:27:51 +00:00
using BotSharp.Abstraction.MLTasks;
2024-04-23 16:25:05 +00:00
using BotSharp.Abstraction.Options;
2024-04-23 19:02:00 +00:00
using BotSharp.Abstraction.Templating;
2024-05-09 20:59:59 +00:00
using BotSharp.Abstraction.Translation.Models;
2024-04-24 04:47:38 +00:00
using System.Collections;
2024-04-23 16:25:05 +00:00
using System.Reflection;
2024-04-22 19:29:21 +00:00
namespace BotSharp.Core.Translation;
public class TranslationService : ITranslationService
{
private readonly IServiceProvider _services;
private readonly ILogger<TranslationService> _logger;
2024-04-23 16:25:05 +00:00
private readonly BotSharpOptions _options;
2024-04-23 19:02:00 +00:00
private Agent _router;
private string _messageId;
2024-04-25 16:27:51 +00:00
private IChatCompletion _completion;
2024-04-22 19:29:21 +00:00
public TranslationService(IServiceProvider services,
2024-04-23 16:25:05 +00:00
ILogger<TranslationService> logger,
BotSharpOptions options)
2024-04-22 19:29:21 +00:00
{
_services = services;
_logger = logger;
2024-04-23 16:25:05 +00:00
_options = options;
2024-04-22 19:29:21 +00:00
}
2024-04-23 19:02:00 +00:00
public async Task<T> Translate<T>(Agent router, string messageId, T data, string language = "Spanish", bool clone = true) where T : class
2024-04-22 19:29:21 +00:00
{
2024-04-23 19:02:00 +00:00
_router = router;
_messageId = messageId;
2024-04-23 22:13:38 +00:00
var unique = new HashSet<string>();
Collect(data, ref unique);
2024-04-25 15:58:28 +00:00
if (unique.IsNullOrEmpty())
2024-04-23 22:13:38 +00:00
{
return data;
}
2024-05-08 16:47:21 +00:00
var clonedData = data;
2024-04-23 16:25:05 +00:00
if (clone)
{
2024-05-08 16:47:21 +00:00
clonedData = Clone(data);
if (clonedData == null)
{
return data;
}
2024-04-23 16:25:05 +00:00
}
2024-04-25 16:27:51 +00:00
// chat completion
_completion = CompletionProvider.GetChatCompletion(_services,
provider: _router?.LlmConfig?.Provider,
model: _router?.LlmConfig?.Model);
var template = _router.Templates.First(x => x.Name == "translation_prompt").Content;
var texts = unique.ToArray();
2024-05-08 16:47:21 +00:00
var translatedStringList = await InnerTranslate(JsonSerializer.Serialize(texts, _options.JsonSerializerOptions), language, template);
2024-04-25 16:27:51 +00:00
try
{
2024-05-09 20:59:59 +00:00
// Override language if it's Unknown, it's used to output the corresponding language.
var states = _services.GetRequiredService<IConversationStateService>();
2024-05-11 02:59:51 +00:00
if (!states.ContainsState("language"))
{
var inputLanguage = string.IsNullOrEmpty(translatedStringList.InputLanguage) ? LanguageType.ENGLISH : translatedStringList.InputLanguage;
states.SetState("language", inputLanguage, activeRounds: 1);
}
2024-05-09 20:59:59 +00:00
var translatedTexts = translatedStringList.Texts;
2024-04-25 16:27:51 +00:00
var map = new Dictionary<string, string>();
for (var i = 0; i < texts.Length; i++)
{
map.Add(texts[i], translatedTexts[i]);
}
2024-05-08 16:47:21 +00:00
clonedData = Assign(clonedData, map);
2024-04-25 16:27:51 +00:00
}
catch (Exception ex)
{
_logger.LogError(ex.Message);
}
2024-04-23 22:13:38 +00:00
2024-05-08 16:47:21 +00:00
return clonedData;
2024-04-23 16:25:05 +00:00
}
private T Clone<T>(T data) where T : class
{
2024-04-23 16:27:20 +00:00
if (data == null) return data;
2024-05-08 16:47:21 +00:00
var str = JsonSerializer.Serialize(data, _options.JsonSerializerOptions);
var cloned = JsonSerializer.Deserialize<T>(str, _options.JsonSerializerOptions);
2024-04-23 16:25:05 +00:00
return cloned;
}
/// <summary>
/// Collect unique strings in data
/// </summary>
/// <typeparam name="T"></typeparam>
/// <param name="data"></param>
/// <param name="res"></param>
private void Collect<T>(T data, ref HashSet<string> res) where T : class
{
if (data == null) return;
var dataType = data.GetType();
2024-04-25 15:58:28 +00:00
if (IsStringType(dataType))
2024-04-23 16:25:05 +00:00
{
res.Add(data.ToString());
return;
}
2024-04-25 15:58:28 +00:00
if (IsDictionaryType(dataType))
2024-04-23 16:25:05 +00:00
{
return;
}
2024-04-25 15:58:28 +00:00
if (IsListType(dataType))
2024-04-23 16:25:05 +00:00
{
var elementType = dataType.IsArray ? dataType.GetElementType() : dataType.GetGenericArguments().FirstOrDefault();
2024-04-25 15:58:28 +00:00
if (IsStringType(elementType))
2024-04-23 16:25:05 +00:00
{
foreach (var item in (data as IEnumerable<string>))
{
if (item == null) continue;
res.Add(item);
}
}
2024-04-25 15:58:28 +00:00
else if (IsTrackToNextLevel(elementType))
2024-04-23 16:25:05 +00:00
{
foreach (var item in (data as IEnumerable<object>))
{
if (item == null) continue;
Collect(item, ref res);
}
}
return;
}
var props = dataType.GetProperties();
foreach (var prop in props)
{
var value = prop.GetValue(data, null);
var propType = prop.PropertyType;
var translate = prop.GetCustomAttributes(true).FirstOrDefault(x => x.GetType() == typeof(TranslateAttribute));
if (value == null) continue;
2024-04-25 15:58:28 +00:00
if (IsStringType(propType))
2024-04-23 16:25:05 +00:00
{
if (translate != null)
{
Collect(value, ref res);
}
}
2024-04-25 15:58:28 +00:00
else if (IsTrackToNextLevel(propType))
2024-04-23 16:25:05 +00:00
{
2024-04-25 15:58:28 +00:00
if (IsDictionaryType(propType))
2024-04-23 16:25:05 +00:00
{
Collect(value, ref res);
}
2024-04-25 15:58:28 +00:00
else if (IsListType(propType))
2024-04-23 16:25:05 +00:00
{
var elementType = propType.IsArray ? propType.GetElementType() : propType.GetGenericArguments().FirstOrDefault();
2024-04-25 15:58:28 +00:00
if (IsStringType(elementType))
2024-04-23 16:25:05 +00:00
{
if (translate != null)
{
Collect(value, ref res);
}
}
2024-04-25 15:58:28 +00:00
else if (IsTrackToNextLevel(elementType))
2024-04-23 16:25:05 +00:00
{
Collect(value, ref res);
}
}
else
{
Collect(value, ref res);
}
}
}
}
/// <summary>
/// Assign translated values to corresponding attributes
/// </summary>
/// <typeparam name="T"></typeparam>
/// <param name="data"></param>
/// <param name="map"></param>
/// <returns></returns>
private T Assign<T>(T data, Dictionary<string, string> map) where T : class
{
if (data == null) return data;
var dataType = data.GetType();
2024-04-25 15:58:28 +00:00
if (IsStringType(dataType) && map.TryGetValue(data.ToString(), out var target))
2024-04-23 16:25:05 +00:00
{
return target as T;
}
2024-04-25 15:58:28 +00:00
if (IsDictionaryType(dataType))
2024-04-23 16:25:05 +00:00
{
return data;
}
2024-04-25 15:58:28 +00:00
if (IsListType(dataType))
2024-04-23 16:25:05 +00:00
{
var elementType = dataType.IsArray ? dataType.GetElementType() : dataType.GetGenericArguments().FirstOrDefault();
2024-04-25 15:58:28 +00:00
if (IsStringType(elementType))
2024-04-23 16:25:05 +00:00
{
var list = new List<string>();
foreach (var item in (data as IEnumerable<string>))
{
if (map.TryGetValue(item, out target))
{
list.Add(target);
}
else
{
list.Add(item?.ToString());
}
}
data = dataType.IsArray ? list.ToArray() as T : list as T;
}
2024-04-25 15:58:28 +00:00
else if (IsTrackToNextLevel(elementType))
2024-04-23 16:25:05 +00:00
{
foreach (var item in (data as IEnumerable<object>))
{
if (item == null) continue;
Assign(item, map);
}
}
return data;
}
var props = dataType.GetProperties();
foreach (var prop in props)
{
var value = prop.GetValue(data, null);
var propType = prop.PropertyType;
var translate = prop.GetCustomAttributes(true).FirstOrDefault(x => x.GetType() == typeof(TranslateAttribute));
if (value == null) continue;
2024-04-25 15:58:28 +00:00
if (IsStringType(propType))
2024-04-23 16:25:05 +00:00
{
if (translate != null)
{
prop.SetValue(data, Assign(value, map));
}
}
2024-04-25 15:58:28 +00:00
else if (IsTrackToNextLevel(propType))
2024-04-23 16:25:05 +00:00
{
2024-04-25 15:58:28 +00:00
if (IsDictionaryType(propType))
2024-04-23 16:25:05 +00:00
{
Assign(value, map);
}
2024-04-25 15:58:28 +00:00
else if (IsListType(propType))
2024-04-23 16:25:05 +00:00
{
var elementType = propType.IsArray ? propType.GetElementType() : propType.GetGenericArguments().FirstOrDefault();
2024-04-25 15:58:28 +00:00
if (IsStringType(elementType))
2024-04-23 16:25:05 +00:00
{
if (translate != null)
{
2024-05-08 16:47:21 +00:00
var json = JsonSerializer.Serialize(Assign(value, map), _options.JsonSerializerOptions);
var targetValue = JsonSerializer.Deserialize(json, propType, _options.JsonSerializerOptions);
2024-04-23 16:25:05 +00:00
prop.SetValue(data, targetValue);
}
}
2024-04-25 15:58:28 +00:00
else if (IsTrackToNextLevel(elementType))
2024-04-23 16:25:05 +00:00
{
prop.SetValue(data, Assign(value, map));
}
}
else
{
Assign(value, map);
}
}
}
2024-04-22 19:29:21 +00:00
return data;
}
2024-04-23 16:25:05 +00:00
/// <summary>
/// Translate
/// </summary>
/// <param name="list"></param>
/// <param name="language"></param>
/// <returns></returns>
2024-05-09 20:59:59 +00:00
private async Task<TranslationOutput> InnerTranslate(string texts, string language, string template)
2024-04-23 16:25:05 +00:00
{
2024-04-23 19:02:00 +00:00
var translator = new Agent
{
Id = Guid.Empty.ToString(),
Name = "Translator",
TemplateDict = new Dictionary<string, object>
{
2024-04-25 16:27:51 +00:00
{ "text_list", texts },
2024-04-23 19:02:00 +00:00
{ "language", language }
}
};
var render = _services.GetRequiredService<ITemplateRender>();
var prompt = render.Render(template, translator.TemplateDict);
var translationDialogs = new List<RoleDialogModel>
{
new RoleDialogModel(AgentRole.User, prompt)
{
FunctionName = "translation_prompt",
MessageId = _messageId
}
};
2024-04-25 16:27:51 +00:00
var response = await _completion.GetChatCompletions(translator, translationDialogs);
2024-05-09 20:59:59 +00:00
return response.Content.JsonContent<TranslationOutput>();
2024-04-23 16:25:05 +00:00
}
2024-04-25 15:58:28 +00:00
#region Type methods
private static bool IsStringType(Type? type)
{
if (type == null) return false;
return type == typeof(string);
}
private static bool IsListType(Type? type)
{
if (type == null) return false;
var interfaces = type.GetTypeInfo().ImplementedInterfaces;
return type.IsArray || interfaces.Any(x => x.Name == typeof(IEnumerable).Name);
}
private static bool IsDictionaryType(Type? type)
{
if (type == null) return false;
var underlyingInterfaces = type.UnderlyingSystemType.GetTypeInfo().ImplementedInterfaces;
return underlyingInterfaces.Any(x => x.Name == typeof(IDictionary).Name);
}
private static bool IsTrackToNextLevel(Type? type)
{
if (type == null) return false;
return type.IsClass || type.IsInterface || type.IsAbstract;
}
#endregion
2024-04-22 19:29:21 +00:00
}