272 lines
7.7 KiB
C#
272 lines
7.7 KiB
C#
using System;
|
|
using System.IO;
|
|
using System.Linq;
|
|
using System.Threading;
|
|
using xAiApi.Constants;
|
|
using xAiModels.Models;
|
|
using xAiApi.Extensions;
|
|
using xAiApi.Interfaces;
|
|
using xCommons.Extensions;
|
|
using xAiModels.Extensions;
|
|
using xAiApi.Configurations;
|
|
using xExceptions.Constants;
|
|
using System.Threading.Tasks;
|
|
using Microsoft.Extensions.AI;
|
|
using System.Collections.Generic;
|
|
using xAiApi.Interfaces.Extractors;
|
|
using System.Runtime.CompilerServices;
|
|
|
|
namespace xAiApi.Providers.Extractors
|
|
{
|
|
/// <summary>
|
|
/// Extracts content from files using Vision Model ...
|
|
/// </summary>
|
|
public class XVisionFileContentExtractor : IXVisionFileContentExtractor
|
|
{
|
|
/// <summary>
|
|
/// Supported MIME Types ...
|
|
/// </summary>
|
|
private static readonly string[] SupportedMimeTypes =
|
|
[
|
|
"image/png",
|
|
"image/jpeg",
|
|
"image/jpg",
|
|
"image/webp",
|
|
"image/bmp",
|
|
"application/pdf",
|
|
];
|
|
|
|
//
|
|
private readonly string prompt;
|
|
private readonly XAiModelDescriptor descriptor;
|
|
|
|
public XVisionFileContentExtractor(
|
|
XAiApiConfiguration configuration,
|
|
string promptName = XAiApiConstants.XAiApiContentExtractionPromptName,
|
|
string modelName = XAiApiConstants.XAiDefaultVisionModelName
|
|
)
|
|
{
|
|
//
|
|
// Retrieve Descriptoir ...
|
|
if (modelName.IsNullOrEmpty())
|
|
{
|
|
XException.InvalidArgs.Throw();
|
|
}
|
|
descriptor = configuration.GetModel(modelName);
|
|
|
|
//
|
|
// Retrieve Extraction Prompt ...
|
|
if (promptName.IsNullOrEmpty())
|
|
{
|
|
XException.InvalidArgs.Throw();
|
|
}
|
|
prompt = configuration.GetPrompt(
|
|
name: promptName,
|
|
@params: null
|
|
);
|
|
|
|
//
|
|
// Validate Requirements ...
|
|
if (prompt.IsNullOrEmpty() || descriptor.IsNullOrDefault())
|
|
{
|
|
XException.InvalidData.Throw();
|
|
}
|
|
}
|
|
|
|
/// <summary>
|
|
/// Check if this extractor supports the specified MIME type ...
|
|
/// </summary>
|
|
public bool CanExtract(string mimeType)
|
|
{
|
|
//
|
|
return SupportedMimeTypes.Contains(
|
|
mimeType?.ToLowerInvariant() ?? string.Empty
|
|
);
|
|
}
|
|
|
|
/// <summary>
|
|
/// Extract text content from file stream ...
|
|
/// </summary>
|
|
public async Task<string> ExtractAsync(
|
|
Stream fileStream,
|
|
string mimeType,
|
|
CancellationToken cancellationToken = default
|
|
)
|
|
{
|
|
var result = await ExtractRichAsync(fileStream, "file", mimeType, cancellationToken);
|
|
return result.Text;
|
|
}
|
|
|
|
/// <summary>
|
|
/// Extract content from stream as Rich Result ...
|
|
/// </summary>
|
|
/// <param name="fileStream"></param>
|
|
/// <param name="fileName"></param>
|
|
/// <param name="mimeType"></param>
|
|
/// <param name="cancellationToken"></param>
|
|
/// <returns></returns>
|
|
public async Task<XFileExtractionResult> ExtractRichAsync(
|
|
Stream fileStream,
|
|
string fileName,
|
|
string mimeType,
|
|
CancellationToken cancellationToken = default
|
|
)
|
|
{
|
|
//
|
|
var result = new XFileExtractionResult
|
|
{
|
|
FileName = fileName,
|
|
MimeType = mimeType
|
|
};
|
|
|
|
//
|
|
try
|
|
{
|
|
//
|
|
using var memoryStream = new MemoryStream();
|
|
await fileStream.CopyToAsync(memoryStream, cancellationToken);
|
|
var imageBytes = memoryStream.ToArray();
|
|
var messages = new[]
|
|
{
|
|
//
|
|
new ChatMessage(ChatRole.User, new[]
|
|
{
|
|
new DataContent(imageBytes, mimeType)
|
|
})
|
|
};
|
|
|
|
//
|
|
var extractedContent = await RequestOCRAsync(
|
|
messages: messages,
|
|
cancellationToken: cancellationToken
|
|
);
|
|
if (!extractedContent.IsNullOrEmpty())
|
|
{
|
|
result.Text = extractedContent.Trim();
|
|
}
|
|
else
|
|
{
|
|
result.ErrorMessage = "there is not any Extracted Content ...";
|
|
}
|
|
}
|
|
catch (Exception ex)
|
|
{
|
|
result.ErrorMessage = $"Error in Vision Extraction: {ex.Message} ...";
|
|
}
|
|
|
|
//
|
|
return result;
|
|
}
|
|
|
|
//
|
|
#region OCR ...
|
|
/// <summary>
|
|
/// Request for Doing OCR on Given Data Contents ...
|
|
/// </summary>
|
|
/// <param name="messages"></param>
|
|
/// <param name="cancellationToken"></param>
|
|
/// <returns></returns>
|
|
public virtual async Task<string> RequestOCRAsync(
|
|
IList<ChatMessage> messages,
|
|
CancellationToken cancellationToken = default
|
|
)
|
|
{
|
|
//
|
|
// Validate ...
|
|
if (!messages.HasChild())
|
|
{
|
|
XException.InvalidArgs.Throw();
|
|
}
|
|
|
|
//
|
|
using var client = descriptor.GetClient();
|
|
|
|
//
|
|
// Preparing Extraction Prompt Message ...
|
|
var pMessage = new ChatMessage(
|
|
ChatRole.System,
|
|
prompt
|
|
);
|
|
|
|
//
|
|
messages = [pMessage, .. messages];
|
|
|
|
//
|
|
var response = await client.GetResponseAsync(
|
|
messages: messages,
|
|
cancellationToken: cancellationToken
|
|
);
|
|
|
|
//
|
|
// Validate Response ...
|
|
if (!response.IsValid())
|
|
{
|
|
//
|
|
// Dispose Client ...
|
|
client.Dispose();
|
|
XException.ActionFailed.Throw();
|
|
}
|
|
|
|
//
|
|
// Retrieve Response Text ...
|
|
var result = response.Text;
|
|
|
|
//
|
|
return result;
|
|
}
|
|
|
|
/// <summary>
|
|
/// Request for Doing OCR on Given Data Contents as Stream ...
|
|
/// </summary>
|
|
/// <param name="messages"></param>
|
|
/// <param name="cancellationToken"></param>
|
|
/// <returns></returns>
|
|
public virtual async IAsyncEnumerable<string> RequestOCRAsEnumerable(
|
|
IList<ChatMessage> messages,
|
|
[EnumeratorCancellation]
|
|
CancellationToken cancellationToken = default
|
|
)
|
|
{
|
|
//
|
|
// Validate ...
|
|
if (!messages.HasChild())
|
|
{
|
|
XException.InvalidArgs.Throw();
|
|
}
|
|
|
|
//
|
|
using var client = descriptor.GetClient();
|
|
|
|
//
|
|
// Preparing Extraction Prompt Message ...
|
|
var pMessage = new ChatMessage(
|
|
ChatRole.System,
|
|
prompt
|
|
);
|
|
|
|
//
|
|
messages = [pMessage, .. messages];
|
|
|
|
//
|
|
var enumerable = client.GetStreamingResponseAsync(
|
|
options: null,
|
|
messages: messages
|
|
);
|
|
|
|
//
|
|
await foreach (var res in enumerable)
|
|
{
|
|
//
|
|
// Cancellation Token ...
|
|
if (cancellationToken.IsCancellationRequested)
|
|
{
|
|
yield break;
|
|
}
|
|
|
|
//
|
|
yield return res.Text;
|
|
}
|
|
}
|
|
#endregion
|
|
}
|
|
} |