Added a prompt injection detection (#857)

Co-authored-by: Thorsten Sommer <SommerEngineering@users.noreply.github.com>
This commit is contained in:
Sabrina-devopsandThorsten Sommer authored and GitHub committed 2026-08-23 11:09:11 +02:00
1 parent d3163badd9
commit 902a01a4d0
55 files changed
+5598 -70

No files matched your search

@@ -8,10 +8,7 @@ using AIStudio.Tools.RAG.RAGProcesses;
namespace AIStudio.Tools.AIJobs;
public sealed class AIJobService(
SettingsManager settingsManager,
MessageBus messageBus,
ILogger<AIJobService> logger)
public sealed class AIJobService(SettingsManager settingsManager, MessageBus messageBus, ILogger<AIJobService> logger)
{
private sealed class AIJobState
{
@@ -24,6 +24,7 @@ public sealed class ContentStreamMetadataJsonConverter : JsonConverter<ContentSt
"Image" => JsonSerializer.Deserialize<ContentStreamImageMetadata?>(rawText, options),
"Document" => JsonSerializer.Deserialize<ContentStreamDocumentMetadata?>(rawText, options),
"Error" => JsonSerializer.Deserialize<ContentStreamErrorMetadata?>(rawText, options),
"PromptInjection" => JsonSerializer.Deserialize<ContentStreamPromptInjectionMetadata?>(rawText, options),
_ => null
};
@@ -10,7 +10,8 @@ namespace AIStudio.Tools;
/// </remarks>
/// <param name="Content">The content to append, or null when this event carries none.</param>
/// <param name="Error">The reported failure, or null when the event was processed successfully.</param>
public readonly record struct ContentStreamProcessedEvent(string? Content, ContentStreamErrorDetails? Error)
/// <param name="PromptInjection">What the runtime filtered out of the content, or null when it filtered nothing.</param>
public readonly record struct ContentStreamProcessedEvent(string? Content, ContentStreamErrorDetails? Error, ContentStreamPromptInjectionDetails? PromptInjection = null)
{
/// <summary>
/// An event which neither produced content nor reported a failure.
@@ -20,4 +21,13 @@ public readonly record struct ContentStreamProcessedEvent(string? Content, Conte
public static ContentStreamProcessedEvent FromContent(string? content) => new(content, null);
public static ContentStreamProcessedEvent FromError(ContentStreamErrorDetails? error) => new(null, error);
/// <summary>
/// An event reporting that suspicious passages were filtered out of the content.
/// </summary>
/// <remarks>
/// Carries no content and no error: the content was delivered by the events before it, and
/// filtering is a notice rather than a failure.
/// </remarks>
public static ContentStreamProcessedEvent FromPromptInjection(ContentStreamPromptInjectionDetails? promptInjection) => new(null, null, promptInjection);
}
@@ -0,0 +1,28 @@
using System.Text.Json.Serialization;
using AIStudio.Tools.Security;
namespace AIStudio.Tools;
// ReSharper disable UnusedAutoPropertyAccessor.Global
// ReSharper disable ClassNeverInstantiated.Global
/// <summary>
/// Reports that the runtime filtered suspected prompt injections out of a file.
/// </summary>
/// <remarks>
/// This is a notice, not a failure: the file was read and everything around the filtered
/// passages is intact. It travels beside the content rather than as an error code, because the
/// app needs the findings themselves to tell the user what was removed.
/// </remarks>
public sealed class ContentStreamPromptInjectionDetails
{
[JsonPropertyName("findings")]
public List<PromptInjectionFinding>? Findings { get; init; }
/// <summary>
/// How many passages were filtered. Can exceed the number of findings, because the runtime
/// caps how many it reports in detail while it filters every single one.
/// </summary>
[JsonPropertyName("redacted_count")]
public int RedactedCount { get; init; }
}
@@ -0,0 +1,11 @@
using System.Text.Json.Serialization;
namespace AIStudio.Tools;
// ReSharper disable UnusedAutoPropertyAccessor.Global
// ReSharper disable ClassNeverInstantiated.Global
public sealed class ContentStreamPromptInjectionMetadata : ContentStreamSseMetadata
{
[JsonPropertyName("PromptInjection")]
public ContentStreamPromptInjectionDetails? PromptInjection { get; init; }
}
@@ -73,6 +73,14 @@ public static class ContentStreamSseHandler
case ContentStreamErrorMetadata errorMetadata:
return ContentStreamProcessedEvent.FromError(errorMetadata.Error);
//
// The runtime filtered suspected prompt injections out of the content. The
// content itself already arrived through the events before this one, so this
// only reports what was removed.
//
case ContentStreamPromptInjectionMetadata promptInjectionMetadata:
return ContentStreamProcessedEvent.FromPromptInjection(promptInjectionMetadata.PromptInjection);
default:
return ContentStreamProcessedEvent.FromContent(sseEvent.Content);
}
+5
View File
@@ -78,6 +78,11 @@ public enum Event
/// </summary>
SHOW_INFO,
/// <summary>
/// Requests display of a prompt-injection alert dialog.
/// </summary>
SHOW_PROMPT_INJECTION_ALERT,
/// <summary>
/// Carries an event received from the Tauri runtime.
/// </summary>
@@ -1,3 +1,5 @@
using AIStudio.Tools.Security;
namespace AIStudio.Tools;
/// <summary>
@@ -17,6 +19,34 @@ namespace AIStudio.Tools;
public readonly record struct FileExtractionResult(FileExtractionOutcome Outcome, string Content, FileExtractionErrorCode ErrorCode, string? ErrorMessage, IReadOnlyList<int> FailedPages, string? DetectedFormat)
{
private static readonly int[] NO_FAILED_PAGES = [];
private static readonly PromptInjectionFinding[] NO_FINDINGS = [];
private readonly IReadOnlyList<PromptInjectionFinding>? promptInjectionFindings;
/// <summary>
/// The prompt-injection attempts the runtime filtered out of the content, if any.
/// </summary>
/// <remarks>
/// This is a notice, not a failure: the passages were removed and the content around them
/// is intact, which is why it does not affect the outcome. The findings exist so the app
/// can tell the user what was removed from their document.
/// </remarks>
public IReadOnlyList<PromptInjectionFinding> PromptInjectionFindings
{
get => this.promptInjectionFindings ?? NO_FINDINGS;
init => this.promptInjectionFindings = value;
}
/// <summary>
/// How many passages were filtered out. May exceed the number of findings, because the
/// runtime caps how many it reports in detail while it filters every single one.
/// </summary>
public int PromptInjectionRedactedCount { get; init; }
/// <summary>
/// Gets a value indicating whether prompt injections were filtered out of the content.
/// </summary>
public bool HasFilteredPromptInjections => this.PromptInjectionRedactedCount > 0;
public static FileExtractionResult Success(string content, string? detectedFormat = null) => new(FileExtractionOutcome.SUCCESS, content, FileExtractionErrorCode.NONE, null, NO_FAILED_PAGES, detectedFormat);
@@ -216,6 +216,9 @@ public sealed class PluginConfiguration(bool isInternal, LuaState state, PluginT
// Config: what should be the start page?
ManagedConfiguration.TryProcessConfiguration(x => x.App, x => x.StartPage, this.Id, settingsTable, dryRun);
// Config: show prompt-injection alert dialogs?
ManagedConfiguration.TryProcessConfiguration(x => x.App, x => x.ShowPromptInjectionAlert, this.Id, settingsTable, dryRun);
// Config: show built-in introduction on the home page?
ManagedConfiguration.TryProcessConfiguration(x => x.App, x => x.ShowIntroduction, this.Id, settingsTable, dryRun);
@@ -1,6 +1,7 @@
using System.Text;
using AIStudio.Chat;
using AIStudio.Tools.Security;
namespace AIStudio.Tools.RAG;
@@ -13,86 +14,117 @@ public static class IRetrievalContextExtensions
sb ??= new StringBuilder();
var index = 0;
//
// One report for the whole retrieval run: a query may pull in dozens of contexts, and
// the user wants to know that something was filtered, not to acknowledge it per context.
//
var guardService = Program.SERVICE_PROVIDER.GetRequiredService<PromptInjectionGuardService>();
await using var reportingScope = guardService.BeginAction();
foreach(var retrievalContext in retrievalContexts)
{
index++;
await retrievalContext.AsMarkdown(sb, index, retrievalContexts.Count, token);
}
return sb.ToString();
}
public static async Task<string> AsMarkdown(this IRetrievalContext retrievalContext, StringBuilder? sb = null, int index = -1, int numTotalRetrievalContexts = -1, CancellationToken token = default)
{
sb ??= new StringBuilder();
var contextBuilder = new StringBuilder();
switch (index)
{
case > 0 when numTotalRetrievalContexts is -1:
sb.AppendLine($"# Retrieval context {index}");
contextBuilder.AppendLine($"# Retrieval context {index}");
break;
case > 0 when numTotalRetrievalContexts > 0:
sb.AppendLine($"# Retrieval context {index} of {numTotalRetrievalContexts}");
contextBuilder.AppendLine($"# Retrieval context {index} of {numTotalRetrievalContexts}");
break;
default:
sb.AppendLine("# Retrieval context");
contextBuilder.AppendLine("# Retrieval context");
break;
}
sb.AppendLine($"Data source name: {retrievalContext.DataSourceName}");
sb.AppendLine($"Content category: {retrievalContext.Category}");
sb.AppendLine($"Content type: {retrievalContext.Type}");
sb.AppendLine($"Content path: {retrievalContext.Path}");
contextBuilder.AppendLine($"Data source name: {retrievalContext.DataSourceName}");
contextBuilder.AppendLine($"Content category: {retrievalContext.Category}");
contextBuilder.AppendLine($"Content type: {retrievalContext.Type}");
contextBuilder.AppendLine($"Content path: {retrievalContext.Path}");
if(retrievalContext.Links.Count > 0)
{
sb.AppendLine("Additional links:");
contextBuilder.AppendLine("Additional links:");
foreach(var link in retrievalContext.Links)
sb.AppendLine($"- {link}");
contextBuilder.AppendLine($"- {link}");
}
var guardService = Program.SERVICE_PROVIDER.GetRequiredService<PromptInjectionGuardService>();
var source = PromptInjectionSource.RetrievalContext(retrievalContext.DataSourceName, retrievalContext.Path);
switch(retrievalContext)
{
case RetrievalTextContext textContext:
sb.AppendLine();
sb.AppendLine("Matched text content:");
sb.AppendLine("````");
sb.AppendLine(textContext.MatchedText);
sb.AppendLine("````");
contextBuilder.AppendLine();
contextBuilder.AppendLine("Matched text content:");
contextBuilder.AppendLine("````");
contextBuilder.AppendLine(textContext.MatchedText);
contextBuilder.AppendLine("````");
if(textContext.SurroundingContent.Count > 0)
{
sb.AppendLine();
sb.AppendLine("Surrounding text content:");
contextBuilder.AppendLine();
contextBuilder.AppendLine("Surrounding text content:");
foreach(var surrounding in textContext.SurroundingContent)
{
sb.AppendLine();
sb.AppendLine("````");
sb.AppendLine(surrounding);
sb.AppendLine("````");
contextBuilder.AppendLine();
contextBuilder.AppendLine("````");
contextBuilder.AppendLine(surrounding);
contextBuilder.AppendLine("````");
}
}
await FilterWhatWeHaveSoFar();
break;
case RetrievalImageContext imageContext:
sb.AppendLine();
sb.AppendLine("Matched image content as base64-encoded data:");
sb.AppendLine("````");
sb.AppendLine(await imageContext.TryAsBase64(token) is (success: true, { } base64Image)
? base64Image
//
// Filtering happens before the image is appended, and only covers the text
// around it. Base64 image data is not prose, and running it through the filter
// would have it treated as one enormous encoded carrier.
//
await FilterWhatWeHaveSoFar();
contextBuilder.AppendLine();
contextBuilder.AppendLine("Matched image content as base64-encoded data:");
contextBuilder.AppendLine("````");
contextBuilder.AppendLine(await imageContext.TryAsBase64(token) is (success: true, { } base64Image)
? base64Image
: string.Empty);
sb.AppendLine("````");
contextBuilder.AppendLine("````");
break;
default:
await FilterWhatWeHaveSoFar();
LOGGER.LogWarning($"The retrieval content type '{retrievalContext.Type}' of data source '{retrievalContext.DataSourceName}' at location '{retrievalContext.Path}' is not supported yet.");
break;
}
sb.AppendLine();
contextBuilder.AppendLine();
sb.Append(contextBuilder);
return sb.ToString();
//
// Replaces what has been built so far with its filtered version. A data source is as
// untrusted as any other external content: it may serve text written to steer the model
// rather than to answer the query.
//
async Task FilterWhatWeHaveSoFar()
{
var sanitized = await guardService.SanitizeAsync(contextBuilder.ToString(), source);
contextBuilder.Clear();
contextBuilder.Append(sanitized);
}
}
}
@@ -0,0 +1,6 @@
using System.Text.Json.Serialization;
namespace AIStudio.Tools.Rust;
/// <param name="Text">The content to filter.</param>
public readonly record struct SanitizePromptInjectionsRequest([property: JsonPropertyName("text")] string Text);
@@ -0,0 +1,13 @@
using System.Text.Json.Serialization;
using AIStudio.Tools.Security;
namespace AIStudio.Tools.Rust;
/// <param name="SanitizedText">The content with the suspicious passages removed. Usable as it stands.</param>
/// <param name="Findings">The passages that were removed, capped by the runtime.</param>
/// <param name="RedactedCount">How many passages were removed in total, which may exceed the number of findings.</param>
public readonly record struct SanitizePromptInjectionsResponse(
[property: JsonPropertyName("sanitized_text")] string SanitizedText,
[property: JsonPropertyName("findings")] IReadOnlyList<PromptInjectionFinding> Findings,
[property: JsonPropertyName("redacted_count")] int RedactedCount);
@@ -0,0 +1,17 @@
namespace AIStudio.Tools.Security;
/// <summary>
/// Asks the UI to tell the user what was filtered out of the content they just used.
/// </summary>
/// <remarks>
/// Carries every result of one user action rather than a single one. Attaching twenty
/// documents at once must produce one dialog listing all of them, not twenty dialogs.
/// </remarks>
/// <param name="Results">What was filtered, per piece of content.</param>
public sealed record PromptInjectionAlertMessage(IReadOnlyList<PromptInjectionScanResult> Results)
{
/// <summary>
/// Gets the total number of filtered passages across all content.
/// </summary>
public int TotalRedactedCount => this.Results.Sum(result => result.RedactedCount);
}
@@ -0,0 +1,31 @@
using System.Text.Json.Serialization;
namespace AIStudio.Tools.Security;
/// <summary>
/// One passage the runtime identified as a prompt-injection attempt and filtered out.
/// </summary>
/// <remarks>
/// The property names are spelled out because the content stream is deserialized without a
/// naming policy, so the names have to match what the runtime sends verbatim.
/// </remarks>
public sealed record PromptInjectionFinding
{
/// <summary>
/// Which rule matched, e.g. "instruction_override".
/// </summary>
[JsonPropertyName("rule_id")]
public string RuleId { get; init; } = string.Empty;
/// <summary>
/// The rule's family, e.g. "exfiltration".
/// </summary>
[JsonPropertyName("category")]
public PromptInjectionFindingCategory Category { get; init; } = PromptInjectionFindingCategory.UNKNOWN;
/// <summary>
/// The passage as it appeared in the content, so the user can see what was removed.
/// </summary>
[JsonPropertyName("snippet")]
public string Snippet { get; init; } = string.Empty;
}
@@ -0,0 +1,19 @@
using System.Text.Json.Serialization;
namespace AIStudio.Tools.Security;
[JsonConverter(typeof(PromptInjectionFindingCategoryJsonConverter))]
public enum PromptInjectionFindingCategory
{
UNKNOWN = 0,
OVERRIDE,
ROLE_OVERRIDE,
EXFILTRATION,
JAILBREAK,
AGENT_MANIPULATION,
DELIMITER_EVASION,
MARKUP_EVASION,
ENCODING_EVASION,
PERSISTENCE,
EVASION,
}
@@ -0,0 +1,23 @@
using AIStudio.Tools.PluginSystem;
namespace AIStudio.Tools.Security;
public static class PromptInjectionFindingCategoryExtensions
{
private static string TB(string fallbackEN) => I18N.I.T(fallbackEN, typeof(PromptInjectionFindingCategoryExtensions).Namespace, nameof(PromptInjectionFindingCategoryExtensions));
public static string GetDisplayName(this PromptInjectionFindingCategory category) => category switch
{
PromptInjectionFindingCategory.OVERRIDE => TB("Attempt to override instructions"),
PromptInjectionFindingCategory.ROLE_OVERRIDE => TB("Attempt to change the AI's role"),
PromptInjectionFindingCategory.EXFILTRATION => TB("Attempt to expose protected data"),
PromptInjectionFindingCategory.JAILBREAK => TB("Attempt to bypass safeguards"),
PromptInjectionFindingCategory.AGENT_MANIPULATION => TB("Attempt to manipulate an agent"),
PromptInjectionFindingCategory.DELIMITER_EVASION => TB("Hidden instructions using delimiters"),
PromptInjectionFindingCategory.MARKUP_EVASION => TB("Hidden instructions using markup"),
PromptInjectionFindingCategory.ENCODING_EVASION => TB("Hidden instructions using encoding"),
PromptInjectionFindingCategory.PERSISTENCE => TB("Persistent or delayed instruction"),
PromptInjectionFindingCategory.EVASION => TB("Obfuscated instruction"),
_ => TB("Unknown"),
};
}
@@ -0,0 +1,51 @@
using System.Text.Json;
using System.Text.Json.Serialization;
namespace AIStudio.Tools.Security;
/// <summary>
/// Reads the finding category in the snake_case spelling the Rust runtime sends.
/// </summary>
/// <remarks>
/// The converter sits on the enum itself because neither path that reads a finding passes
/// JsonSerializerOptions: the sanitize response is read by RustService.SanitizePromptInjections
/// and the content stream by RustService.ReadFileContent. The shared RustEnumConverter therefore
/// never applies here, and without a converter on the type only numbers would be accepted.
///
/// An unrecognized category falls back to UNKNOWN instead of throwing. Throwing would cost more
/// than the label: it fails the whole response, and the guard service then passes the content
/// through unfiltered rather than losing a single name.
/// </remarks>
public sealed class PromptInjectionFindingCategoryJsonConverter : JsonConverter<PromptInjectionFindingCategory>
{
private static readonly ILogger<PromptInjectionFindingCategoryJsonConverter> LOG = Program.LOGGER_FACTORY.CreateLogger<PromptInjectionFindingCategoryJsonConverter>();
public override PromptInjectionFindingCategory Read(ref Utf8JsonReader reader, Type typeToConvert, JsonSerializerOptions options)
{
if (reader.TokenType is not JsonTokenType.String)
{
LOG.LogWarning("Cannot read a prompt injection finding category from a '{TokenType}' token. Using UNKNOWN.", reader.TokenType);
return PromptInjectionFindingCategory.UNKNOWN;
}
var text = reader.GetString();
if (string.IsNullOrWhiteSpace(text))
{
LOG.LogWarning("Read an empty prompt injection finding category. Using UNKNOWN.");
return PromptInjectionFindingCategory.UNKNOWN;
}
//
// The enum members are the wire value in upper case, so upper-casing replaces a naming
// policy. Values starting with a digit or sign are rejected up front, because Enum.TryParse
// would otherwise accept "0" or "-1" as a category:
//
if (!char.IsAsciiDigit(text[0]) && text[0] is not ('-' or '+') && Enum.TryParse<PromptInjectionFindingCategory>(text.ToUpperInvariant(), out var category))
return category;
LOG.LogWarning("The runtime reported the unknown prompt injection finding category '{Category}'. Using UNKNOWN.", text);
return PromptInjectionFindingCategory.UNKNOWN;
}
public override void Write(Utf8JsonWriter writer, PromptInjectionFindingCategory value, JsonSerializerOptions options) => writer.WriteStringValue(value.ToString().ToLowerInvariant());
}
@@ -0,0 +1,159 @@
using AIStudio.Settings;
using AIStudio.Tools.PluginSystem;
using AIStudio.Tools.Services;
namespace AIStudio.Tools.Security;
/// <summary>
/// Filters prompt injections out of external content before it reaches a model.
/// </summary>
/// <remarks>
/// The detection itself lives in the Rust runtime. File content is filtered while the runtime
/// streams it, so it never passes through here; what this service adds is the path for content
/// the runtime does not read itself — web pages and retrieval contexts — and the reporting the
/// user sees.
/// </remarks>
public sealed class PromptInjectionGuardService(
RustService rustService,
SettingsManager settingsManager,
ILogger<PromptInjectionGuardService> logger,
ILoggerFactory loggerFactory)
{
public const string WIKI_URL = "https://en.wikipedia.org/wiki/Prompt_engineering#Prompt_injection";
private const string DETECTION_LOG_CATEGORY = "PromptInjectionProtection";
private static string TB(string fallbackEN) => I18N.I.T(fallbackEN, typeof(PromptInjectionGuardService).Namespace, nameof(PromptInjectionGuardService));
private readonly ILogger detectionLogger = loggerFactory.CreateLogger(DETECTION_LOG_CATEGORY);
private readonly Lock reportLock = new();
private readonly List<PromptInjectionScanResult> pendingResults = [];
private int openActions;
/// <summary>
/// Filters prompt injections out of a text the runtime did not read itself, such as a web
/// page or a retrieval context.
/// </summary>
/// <remarks>
/// Returns usable text in every case. When the runtime cannot be reached, the text is passed
/// through unchanged: refusing the user's content because a check could not run would cost
/// them their work over a check that is best-effort anyway. The failure is logged and shown,
/// so it does not pass silently.
/// </remarks>
/// <param name="text">The content to filter.</param>
/// <param name="source">Where the content came from, for the report shown to the user.</param>
/// <returns>The content with any suspicious passages removed.</returns>
public async Task<string> SanitizeAsync(string text, PromptInjectionSource source)
{
if (string.IsNullOrWhiteSpace(text))
return text;
if (await rustService.SanitizePromptInjections(text) is not { } response)
{
logger.LogError("Could not check {SourceKind} '{SourceLabel}' for prompt injections. The content is used unchanged.", source.Kind, source.Label);
await MessageBus.INSTANCE.SendWarning(new(
Icons.Material.Filled.GppMaybe,
string.Format(TB("AI Studio could not check '{0}' for prompt injections. The content is used as it is."), source.NotificationLabel)));
return text;
}
if (response.RedactedCount > 0)
await this.ReportAsync(new(source, response.Findings, response.RedactedCount));
return response.SanitizedText;
}
/// <summary>
/// Records what was filtered out of one piece of content and tells the user about it.
/// </summary>
/// <remarks>
/// Within a BeginAction scope the result is collected and reported together
/// with the rest of that action. Outside of one it is reported immediately: a result that
/// simply waited for the next scope would either never reach the user, or reach them as
/// part of an unrelated action later on.
/// </remarks>
public async Task ReportAsync(PromptInjectionScanResult result)
{
if (!result.WasFiltered)
return;
bool reportNow;
lock (this.reportLock)
{
this.pendingResults.Add(result);
reportNow = this.openActions is 0;
}
if (reportNow)
await this.ReportPendingAsync();
}
/// <summary>
/// Marks the start of one user action, such as attaching a batch of files or sending a
/// message.
/// </summary>
/// <remarks>
/// Results are collected until the action finishes, so the user gets one report about
/// twenty documents instead of twenty reports. Actions may nest: only the outermost one
/// reports.
/// </remarks>
/// <returns>A scope that reports what was filtered once it is disposed.</returns>
public ReportingScope BeginAction()
{
lock (this.reportLock)
this.openActions++;
return new(this);
}
private async Task EndActionAsync()
{
lock (this.reportLock)
{
this.openActions--;
// An inner scope reports nothing: the action the user started is still running.
if (this.openActions > 0)
return;
}
await this.ReportPendingAsync();
}
private async Task ReportPendingAsync()
{
List<PromptInjectionScanResult> results;
lock (this.reportLock)
{
if (this.pendingResults.Count is 0)
return;
results = [..this.pendingResults];
this.pendingResults.Clear();
}
var totalCount = results.Sum(result => result.RedactedCount);
this.detectionLogger.LogWarning(
"Detected and removed {PassageCount} potentially dangerous passage(s) in {SourceCount} content source(s).",
totalCount,
results.Count);
await MessageBus.INSTANCE.SendWarning(new(
Icons.Material.Filled.GppMaybe,
results.Count is 1
? string.Format(TB("AI Studio removed suspicious instructions from '{0}' before using it."), results[0].Source.NotificationLabel)
: string.Format(TB("AI Studio removed suspicious instructions from {0} sources before using them."), results.Count)));
if (settingsManager.ConfigurationData.App.ShowPromptInjectionAlert)
await MessageBus.INSTANCE.SendMessage<PromptInjectionAlertMessage>(null, Event.SHOW_PROMPT_INJECTION_ALERT, new(results));
}
/// <summary>
/// Reports everything filtered during one user action when it goes out of scope.
/// </summary>
public sealed class ReportingScope(PromptInjectionGuardService guardService) : IAsyncDisposable
{
public async ValueTask DisposeAsync() => await guardService.EndActionAsync();
}
}
@@ -0,0 +1,19 @@
namespace AIStudio.Tools.Security;
/// <summary>
/// What the runtime filtered out of one piece of external content.
/// </summary>
/// <param name="Source">Where the content came from, so the user can tell which file or page it was.</param>
/// <param name="Findings">The passages that were removed. Capped by the runtime.</param>
/// <param name="RedactedCount">How many passages were removed in total, which may exceed the number of findings.</param>
public sealed record PromptInjectionScanResult(PromptInjectionSource Source, IReadOnlyList<PromptInjectionFinding> Findings, int RedactedCount)
{
/// <summary>
/// Gets a value indicating whether anything was filtered out of this content.
/// </summary>
/// <remarks>
/// The content itself stays usable either way: passages are removed, the content around
/// them is not rejected.
/// </remarks>
public bool WasFiltered => this.RedactedCount > 0;
}
@@ -0,0 +1,16 @@
namespace AIStudio.Tools.Security;
public readonly record struct PromptInjectionSource(PromptInjectionSourceKind Kind, string Label)
{
public string NotificationLabel => this.Kind is PromptInjectionSourceKind.FILE_CONTENT or PromptInjectionSourceKind.CHAT_ATTACHMENT
? Path.GetFileName(this.Label)
: this.Label;
public static PromptInjectionSource WebContent(string url) => new(PromptInjectionSourceKind.WEB_CONTENT, url);
public static PromptInjectionSource FileContent(string filePath) => new(PromptInjectionSourceKind.FILE_CONTENT, filePath);
public static PromptInjectionSource ChatAttachment(string filePath) => new(PromptInjectionSourceKind.CHAT_ATTACHMENT, filePath);
public static PromptInjectionSource RetrievalContext(string dataSourceName, string path) => new(PromptInjectionSourceKind.RETRIEVAL_CONTEXT, $"{dataSourceName}: {path}");
}
@@ -0,0 +1,10 @@
namespace AIStudio.Tools.Security;
public enum PromptInjectionSourceKind
{
UNKNOWN = 0,
WEB_CONTENT,
FILE_CONTENT,
CHAT_ATTACHMENT,
RETRIEVAL_CONTEXT,
}
@@ -0,0 +1,17 @@
using AIStudio.Tools.PluginSystem;
namespace AIStudio.Tools.Security;
public static class PromptInjectionSourceKindExtensions
{
private static string TB(string fallbackEN) => I18N.I.T(fallbackEN, typeof(PromptInjectionSourceKindExtensions).Namespace, nameof(PromptInjectionSourceKindExtensions));
public static string GetDisplayName(this PromptInjectionSourceKind kind) => kind switch
{
PromptInjectionSourceKind.WEB_CONTENT => TB("Web content"),
PromptInjectionSourceKind.FILE_CONTENT => TB("File content"),
PromptInjectionSourceKind.CHAT_ATTACHMENT => TB("Chat attachment"),
PromptInjectionSourceKind.RETRIEVAL_CONTEXT => TB("Retrieved context"),
_ => TB("Unknown"),
};
}
@@ -1,5 +1,6 @@
using System.Text;
using System.Text.Json;
using AIStudio.Tools.Security;
namespace AIStudio.Tools.Services;
@@ -17,6 +18,13 @@ public sealed partial class RustService
public async Task<FileExtractionResult> ReadArbitraryFileData(string path, int maxChunks, bool extractImages = false)
{
//
// The runtime filters prompt injections while it streams the file. Doing it there rather
// than here means the whole document never has to exist in memory at once, which is what
// makes documents of a few thousand pages affordable.
//
var guardService = Program.SERVICE_PROVIDER.GetRequiredService<PromptInjectionGuardService>();
var streamId = Guid.NewGuid().ToString();
var requestUri = $"/retrieval/fs/extract?path={Uri.EscapeDataString(path)}&stream_id={streamId}&extract_images={extractImages}";
@@ -25,6 +33,8 @@ public sealed partial class RustService
var resultBuilder = new StringBuilder();
var failedPages = new List<int>();
var promptInjectionFindings = new List<PromptInjectionFinding>();
var promptInjectionRedactedCount = 0;
var hasPartialFailure = false;
var failureCode = FileExtractionErrorCode.NONE;
string? failureMessage = null;
@@ -124,6 +134,17 @@ public sealed partial class RustService
detectedFormat = error.DetectedFormat;
}
}
else if (processedEvent.PromptInjection is { } promptInjection)
{
//
// Not a failure: the passages were removed and the document around them is
// intact. It only needs to reach the user, so they know their document was
// changed before the AI saw it.
//
promptInjectionRedactedCount += promptInjection.RedactedCount;
if (promptInjection.Findings is { } findings)
promptInjectionFindings.AddRange(findings);
}
else if (processedEvent.Content is not null)
resultBuilder.AppendLine(processedEvent.Content);
@@ -174,8 +195,28 @@ public sealed partial class RustService
return FileExtractionResult.Failed(FileExtractionErrorCode.NO_CONTENT, "Reading the file produced no content.");
}
return hasPartialFailure
var result = hasPartialFailure
? FileExtractionResult.Partial(content, failedPages, detectedFormat)
: FileExtractionResult.Success(content, detectedFormat);
if (promptInjectionRedactedCount is 0)
return result;
//
// Reported from here rather than from the callers: every way of reading a file passes
// through this method, so this is the one place where no caller can forget it.
//
await guardService.ReportAsync(new(PromptInjectionSource.FileContent(path), promptInjectionFindings, promptInjectionRedactedCount));
//
// Filtering does not change the outcome: the passages were removed and the document
// around them is intact. The findings travel along so a caller can show them next to
// the document they belong to.
//
return result with
{
PromptInjectionFindings = promptInjectionFindings,
PromptInjectionRedactedCount = promptInjectionRedactedCount,
};
}
}
@@ -0,0 +1,49 @@
using AIStudio.Tools.Rust;
namespace AIStudio.Tools.Services;
public sealed partial class RustService
{
/// <summary>
/// How long one sanitize request may take.
/// </summary>
/// <remarks>
/// Web pages and retrieval contexts are small, so this only exists to keep a stuck runtime
/// from blocking the caller forever.
/// </remarks>
private static readonly TimeSpan SANITIZE_TIMEOUT = TimeSpan.FromSeconds(30);
/// <summary>
/// Asks the runtime to filter prompt injections out of a text.
/// </summary>
/// <remarks>
/// File content does not go through here: the runtime filters it while it streams the file.
/// This is the path for content the app fetched itself, i.e. web pages and retrieval contexts.
/// </remarks>
/// <param name="text">The content to filter.</param>
/// <returns>The filtered content and what was found or null when the runtime could not be reached.</returns>
public async Task<SanitizePromptInjectionsResponse?> SanitizePromptInjections(string text)
{
try
{
using var timeoutTokenSource = new CancellationTokenSource(SANITIZE_TIMEOUT);
using var response = await this.http.PostAsJsonAsync(
"/security/prompt-injection/sanitize",
new SanitizePromptInjectionsRequest(text),
cancellationToken: timeoutTokenSource.Token);
if (!response.IsSuccessStatusCode)
{
this.logger?.LogError("Failed to check a text for prompt injections. Status: {StatusCode}, reason: '{ReasonPhrase}'", response.StatusCode, response.ReasonPhrase);
return null;
}
return await response.Content.ReadFromJsonAsync<SanitizePromptInjectionsResponse>(timeoutTokenSource.Token);
}
catch (Exception exception)
{
this.logger?.LogError(exception, "Failed to check a text for prompt injections.");
return null;
}
}
}