made tokenizer concurrent and threadsafe

This commit is contained in:
Paul Koudelka committed 2026-08-10 21:36:01 +02:00
1 parent 6985678d2f
commit c8c1e38475
11 files changed
+110 -306

No files matched your search

@@ -1259,12 +1259,12 @@ public partial class ChatComponent : MSGComponentBase, IAsyncDisposable
return;
}
var response = await this.RustService.GetTokenCount(this.Provider.InstanceName, this.Provider.TokenizerPath, this.inputField.Value);
var response = await this.RustService.GetTokenCount(this.Provider, this.inputField.Value);
if (response is null)
return;
if (!response.Value.Success)
{
this.Logger.LogWarning($"Failed to calculate token count: status='{response.Value.Status}', reason='{response.Value.Message}'");
this.Logger.LogWarning("Failed to calculate token count: reason='{Reason}'", response.Value.Message);
return;
}
this.tokenCount = response.Value.TokenCount.ToString();
@@ -4,5 +4,4 @@ public readonly record struct TokenizerResponse(
bool Success,
int TokenCount,
string Message,
TokenizerStatus Status = TokenizerStatus.UNAVAILABLE,
string StoredPath = "");
@@ -1,8 +0,0 @@
namespace AIStudio.Tools.Rust;
public enum TokenizerStatus
{
UNAVAILABLE,
RUNNING,
AVAILABLE,
}
@@ -66,7 +66,7 @@ public sealed partial class DataSourceEmbeddingService
{
var segments = new List<ExtractedFileSegment>();
await foreach (var segment in rustService.StreamArbitraryFileDataWithTokenCounts(filePath, embeddingProvider.Name, embeddingProvider.TokenizerPath, token))
await foreach (var segment in rustService.StreamArbitraryFileDataWithTokenCounts(filePath, embeddingProvider, token))
{
var normalized = NormalizeChunkSegment(segment.Content);
if (!string.IsNullOrWhiteSpace(normalized))
@@ -536,8 +536,8 @@ public sealed partial class DataSourceEmbeddingService
private async Task<int> GetEmbeddingTokenCountAsync(EmbeddingProvider embeddingProvider, string text, CancellationToken token)
{
var response = await rustService.GetTokenCount(embeddingProvider.Name, embeddingProvider.TokenizerPath, text, token);
if (response is { Success: true, Status: TokenizerStatus.AVAILABLE })
var response = await rustService.GetTokenCount(embeddingProvider, text, token);
if (response is { Success: true })
return response.Value.TokenCount;
var message = response?.Message ?? "No response was returned by the tokenizer service.";
@@ -6,7 +6,6 @@ using AIStudio.Tools.Databases;
using AIStudio.Tools.Databases.EmbeddingState;
using AIStudio.Tools.Databases.VectorStore;
using AIStudio.Tools.RAG;
using AIStudio.Tools.Rust;
namespace AIStudio.Tools.Services;
@@ -160,12 +159,8 @@ public sealed class DataSourceLocalRetrievalService(
return false;
}
var tokenCountResponse = await rustService.GetTokenCount(
embeddingProvider.Name,
embeddingProvider.TokenizerPath,
query,
token);
if (tokenCountResponse is not { Success: true, Status: TokenizerStatus.AVAILABLE })
var tokenCountResponse = await rustService.GetTokenCount(embeddingProvider, query, token);
if (tokenCountResponse is not { Success: true })
{
logger.LogWarning(
"Skipping vector retrieval for data source '{DataSourceName}' ({DataSourceId}) because the token count for embedding provider '{EmbeddingProviderName}' could not be determined. Reason='{Reason}'.",
@@ -2,6 +2,7 @@ using System.Text;
using System.Text.Json;
using System.Runtime.CompilerServices;
using AIStudio.Settings;
using AIStudio.Tools.Rust;
namespace AIStudio.Tools.Services;
@@ -85,37 +86,21 @@ public sealed partial class RustService
public async IAsyncEnumerable<string> StreamArbitraryFileData(string path, bool extractImages = false, [EnumeratorCancellation] CancellationToken token = default)
{
await foreach (var segment in this.StreamArbitraryFileDataCore(path, extractImages, false, token))
await foreach (var segment in this.StreamArbitraryFileDataCore(path, extractImages, false, string.Empty, token))
yield return segment.Content;
}
public async IAsyncEnumerable<ArbitraryFileDataSegment> StreamArbitraryFileDataWithTokenCounts(
string path,
string providerName,
string tokenizerPath,
EmbeddingProvider embeddingProvider,
[EnumeratorCancellation] CancellationToken token = default)
{
await this.tokenizerLock.WaitAsync(token);
try
await foreach (var segment in this.StreamArbitraryFileDataCore(path, false, true, embeddingProvider.TokenizerPath, token))
{
var tokenizerResponse = await this.EnsureTokenizerCoreAsync(providerName, tokenizerPath);
if (tokenizerResponse is not { Success: true, Status: TokenizerStatus.AVAILABLE })
{
var message = tokenizerResponse?.Message ?? "No response was returned by the tokenizer service.";
throw new InvalidOperationException($"Could not initialize tokenizer for provider '{providerName}'. {message}");
}
if (segment.TokenCount is null)
throw new InvalidOperationException($"Rust did not return a token count for an extracted segment from '{path}' using provider '{embeddingProvider.Name}'.");
await foreach (var segment in this.StreamArbitraryFileDataCore(path, false, true, token))
{
if (segment.TokenCount is null)
throw new InvalidOperationException($"Rust did not return a token count for an extracted segment from '{path}'.");
yield return new(segment.Content, segment.TokenCount.Value);
}
}
finally
{
this.tokenizerLock.Release();
yield return new(segment.Content, segment.TokenCount.Value);
}
}
@@ -123,10 +108,11 @@ public sealed partial class RustService
string path,
bool extractImages,
bool includeTokenCount,
string tokenizerPath,
[EnumeratorCancellation] CancellationToken token)
{
var streamId = Guid.NewGuid().ToString();
var requestUri = $"/retrieval/fs/extract?path={Uri.EscapeDataString(path)}&stream_id={streamId}&extract_images={extractImages}&include_token_count={includeTokenCount}";
var requestUri = $"/retrieval/fs/extract?path={Uri.EscapeDataString(path)}&stream_id={streamId}&extract_images={extractImages}&include_token_count={includeTokenCount}&tokenizer_path={Uri.EscapeDataString(tokenizerPath)}";
using var request = new HttpRequestMessage(HttpMethod.Get, requestUri);
using var response = await this.http.SendAsync(request, HttpCompletionOption.ResponseHeadersRead, token);
@@ -1,4 +1,4 @@
using AIStudio.Provider;
using AIStudio.Settings;
using AIStudio.Tools.Rust;
namespace AIStudio.Tools.Services;
@@ -7,35 +7,12 @@ public sealed partial class RustService
{
internal const int MAX_TOKEN_COUNT_REQUEST_TEXT_LENGTH = 200_000;
private readonly SemaphoreSlim tokenizerLock = new(1, 1);
private string currentTokenizerPath = string.Empty;
private bool hasInitializedTokenizer;
private static TokenizerResponse CreateUnavailableTokenizerResponse(string message) => new(
false,
0,
message,
TokenizerStatus.UNAVAILABLE,
string.Empty);
public async Task<TokenizerResponse> GetTokenizerInfo(CancellationToken cancellationToken = default)
{
try
{
return await this.http.GetFromJsonAsync<TokenizerResponse>("/system/tokenizer/info", this.jsonRustSerializerOptions, cancellationToken);
}
catch (OperationCanceledException) when (cancellationToken.IsCancellationRequested)
{
this.logger?.LogWarning("Fetching tokenizer info from Rust service was cancelled by caller.");
return CreateUnavailableTokenizerResponse("Operation cancelled by caller.");
}
catch (Exception e)
{
this.logger?.LogError(e, "Error while fetching tokenizer info from Rust service.");
return CreateUnavailableTokenizerResponse(e.Message);
}
}
public async Task<TokenizerResponse> ValidateTokenizer(string filePath)
{
var result = await this.http.PostAsJsonAsync("/tokenizer/validate", new {
@@ -86,99 +63,25 @@ public sealed partial class RustService
return await result.Content.ReadFromJsonAsync<TokenizerResponse>(this.jsonRustSerializerOptions);
}
public Task<TokenizerResponse?> GetTokenCount(string text)
{
return this.GetTokenCountCoreAsync(text);
}
public Task<TokenizerResponse?> GetTokenCount(AIStudio.Settings.Provider provider, string text, CancellationToken cancellationToken = default) =>
this.GetTokenCount(provider.InstanceName, provider.TokenizerPath, text, cancellationToken);
public async Task<TokenizerResponse?> GetTokenCount(string providerName, string path, string text, CancellationToken cancellationToken = default)
{
await this.tokenizerLock.WaitAsync(cancellationToken);
try
{
var tokenizerResponse = await this.EnsureTokenizerCoreAsync(providerName, path);
if (tokenizerResponse is not { Success: true, Status: TokenizerStatus.AVAILABLE })
return tokenizerResponse;
public Task<TokenizerResponse?> GetTokenCount(EmbeddingProvider provider, string text, CancellationToken cancellationToken = default) =>
this.GetTokenCount(provider.Name, provider.TokenizerPath, text, cancellationToken);
return await this.GetTokenCountCoreAsync(text, cancellationToken);
}
finally
{
this.tokenizerLock.Release();
}
}
private async Task<TokenizerResponse?> GetTokenCountCoreAsync(string text, CancellationToken cancellationToken = default)
private async Task<TokenizerResponse?> GetTokenCount(string providerName, string tokenizerPath, string text, CancellationToken cancellationToken)
{
var result = await this.http.PostAsJsonAsync("/tokenizer/count", new {
text = text,
tokenizer_path = tokenizerPath,
}, this.jsonRustSerializerOptions, cancellationToken);
if (!result.IsSuccessStatusCode)
{
this.logger!.LogError($"Failed to get the token count '{result.StatusCode}'");
this.hasInitializedTokenizer = false;
this.logger!.LogError("Failed to get the token count for provider '{ProviderName}': {StatusCode}", providerName, result.StatusCode);
return CreateUnavailableTokenizerResponse("Error while getting token count from Rust service: "+result.StatusCode);
}
var response = await result.Content.ReadFromJsonAsync<TokenizerResponse>(this.jsonRustSerializerOptions);
if (response is not { Status: TokenizerStatus.AVAILABLE })
this.hasInitializedTokenizer = false;
return response;
}
public async Task<TokenizerResponse?> SetTokenizer(string providerName, string path)
{
this.logger!.LogInformation($"Setting a new tokenizer for '{providerName}'");
var result = await this.http.PostAsJsonAsync("/tokenizer/set", new {
file_path = path,
}, this.jsonRustSerializerOptions);
if (!result.IsSuccessStatusCode)
{
this.logger!.LogError($"Failed to set the tokenizer '{result.StatusCode}'");
this.hasInitializedTokenizer = false;
return CreateUnavailableTokenizerResponse("An error occured while sending the path to the Rust framework for setting a tokenizer: "+result.StatusCode);
}
var response = await result.Content.ReadFromJsonAsync<TokenizerResponse>(this.jsonRustSerializerOptions);
if (response is not { Success: true, Status: TokenizerStatus.AVAILABLE })
this.hasInitializedTokenizer = false;
return response;
}
public async Task<TokenizerResponse?> EnsureTokenizer(string providerName, string path)
{
await this.tokenizerLock.WaitAsync();
try
{
return await this.EnsureTokenizerCoreAsync(providerName, path);
}
finally
{
this.tokenizerLock.Release();
}
}
private async Task<TokenizerResponse?> EnsureTokenizerCoreAsync(string providerName, string path)
{
if (this.hasInitializedTokenizer && this.currentTokenizerPath == path)
return new TokenizerResponse(true, 0, string.Empty, TokenizerStatus.AVAILABLE);
var response = await this.SetTokenizer(providerName, path);
if (response is { Success: true, Status: TokenizerStatus.AVAILABLE })
{
this.currentTokenizerPath = path;
this.hasInitializedTokenizer = true;
}
else
{
this.currentTokenizerPath = string.Empty;
this.hasInitializedTokenizer = false;
}
return response;
return await result.Content.ReadFromJsonAsync<TokenizerResponse>(this.jsonRustSerializerOptions, cancellationToken);
}
}
@@ -94,7 +94,6 @@ public sealed partial class RustService : BackgroundService
{
this.http.Dispose();
this.userLanguageLock.Dispose();
this.tokenizerLock.Dispose();
this.userNameLock.Dispose();
base.Dispose();
}