Update prediction interface to provide additional feedback to a predictor plugin (#15421)

This commit is contained in:
Dongbo Wang
2021-05-20 15:52:13 -07:00
committed by GitHub
parent e927e94aed
commit 5febcad3db
5 changed files with 270 additions and 91 deletions
@@ -10,7 +10,7 @@ using System.Management.Automation.Language;
using System.Threading;
using System.Threading.Tasks;
namespace System.Management.Automation.Subsystem
namespace System.Management.Automation.Subsystem.Prediction
{
/// <summary>
/// The class represents the prediction result from a predictor.
@@ -59,9 +59,9 @@ namespace System.Management.Automation.Subsystem
/// <param name="ast">The <see cref="Ast"/> object from parsing the current command line input.</param>
/// <param name="astTokens">The <see cref="Token"/> objects from parsing the current command line input.</param>
/// <returns>A list of <see cref="PredictionResult"/> objects.</returns>
public static Task<List<PredictionResult>?> PredictInput(string client, Ast ast, Token[] astTokens)
public static Task<List<PredictionResult>?> PredictInputAsync(PredictionClient client, Ast ast, Token[] astTokens)
{
return PredictInput(client, ast, astTokens, millisecondsTimeout: 20);
return PredictInputAsync(client, ast, astTokens, millisecondsTimeout: 20);
}
/// <summary>
@@ -72,7 +72,7 @@ namespace System.Management.Automation.Subsystem
/// <param name="astTokens">The <see cref="Token"/> objects from parsing the current command line input.</param>
/// <param name="millisecondsTimeout">The milliseconds to timeout.</param>
/// <returns>A list of <see cref="PredictionResult"/> objects.</returns>
public static async Task<List<PredictionResult>?> PredictInput(string client, Ast ast, Token[] astTokens, int millisecondsTimeout)
public static async Task<List<PredictionResult>?> PredictInputAsync(PredictionClient client, Ast ast, Token[] astTokens, int millisecondsTimeout)
{
Requires.Condition(millisecondsTimeout > 0, nameof(millisecondsTimeout));
@@ -86,17 +86,13 @@ namespace System.Management.Automation.Subsystem
var tasks = new Task<PredictionResult?>[predictors.Count];
using var cancellationSource = new CancellationTokenSource();
Func<object?, PredictionResult?> callBack = GetCallBack(client, context, cancellationSource);
for (int i = 0; i < predictors.Count; i++)
{
ICommandPredictor predictor = predictors[i];
tasks[i] = Task.Factory.StartNew(
state =>
{
var predictor = (ICommandPredictor)state!;
SuggestionPackage pkg = predictor.GetSuggestion(client, context, cancellationSource.Token);
return pkg.SuggestionEntries?.Count > 0 ? new PredictionResult(predictor.Id, predictor.Name, pkg.Session, pkg.SuggestionEntries) : null;
},
callBack,
predictor,
cancellationSource.Token,
TaskCreationOptions.DenyChildAttach,
@@ -122,6 +118,21 @@ namespace System.Management.Automation.Subsystem
}
return resultList;
// A local helper function to avoid creating an instance of the generated delegate helper class
// when no predictor is registered.
static Func<object?, PredictionResult?> GetCallBack(
PredictionClient client,
PredictionContext context,
CancellationTokenSource cancellationSource)
{
return state =>
{
var predictor = (ICommandPredictor)state!;
SuggestionPackage pkg = predictor.GetSuggestion(client, context, cancellationSource.Token);
return pkg.SuggestionEntries?.Count > 0 ? new PredictionResult(predictor.Id, predictor.Name, pkg.Session, pkg.SuggestionEntries) : null;
};
}
}
/// <summary>
@@ -129,7 +140,7 @@ namespace System.Management.Automation.Subsystem
/// </summary>
/// <param name="client">Represents the client that initiates the call.</param>
/// <param name="history">History command lines provided as references for prediction.</param>
public static void OnCommandLineAccepted(string client, IReadOnlyList<string> history)
public static void OnCommandLineAccepted(PredictionClient client, IReadOnlyList<string> history)
{
Requires.NotNull(history, nameof(history));
@@ -139,16 +150,54 @@ namespace System.Management.Automation.Subsystem
return;
}
Action<ICommandPredictor>? callBack = null;
foreach (ICommandPredictor predictor in predictors)
{
if (predictor.SupportEarlyProcessing)
if (predictor.CanAcceptFeedback(client, PredictorFeedbackKind.CommandLineAccepted))
{
ThreadPool.QueueUserWorkItem<ICommandPredictor>(
state => state.StartEarlyProcessing(client, history),
predictor,
preferLocal: false);
callBack ??= GetCallBack(client, history);
ThreadPool.QueueUserWorkItem<ICommandPredictor>(callBack, predictor, preferLocal: false);
}
}
// A local helper function to avoid creating an instance of the generated delegate helper class
// when no predictor is registered, or no registered predictor accepts this feedback.
static Action<ICommandPredictor> GetCallBack(PredictionClient client, IReadOnlyList<string> history)
{
return predictor => predictor.OnCommandLineAccepted(client, history);
}
}
/// <summary>
/// Allow registered predictors to know the execution result (success/failure) of the last accepted command line.
/// </summary>
/// <param name="client">Represents the client that initiates the call.</param>
/// <param name="commandLine">The last accepted command line.</param>
/// <param name="success">Whether the execution of the last command line was successful.</param>
public static void OnCommandLineExecuted(PredictionClient client, string commandLine, bool success)
{
var predictors = SubsystemManager.GetSubsystems<ICommandPredictor>();
if (predictors.Count == 0)
{
return;
}
Action<ICommandPredictor>? callBack = null;
foreach (ICommandPredictor predictor in predictors)
{
if (predictor.CanAcceptFeedback(client, PredictorFeedbackKind.CommandLineExecuted))
{
callBack ??= GetCallBack(client, commandLine, success);
ThreadPool.QueueUserWorkItem<ICommandPredictor>(callBack, predictor, preferLocal: false);
}
}
// A local helper function to avoid creating an instance of the generated delegate helper class
// when no predictor is registered, or no registered predictor accepts this feedback.
static Action<ICommandPredictor> GetCallBack(PredictionClient client, string commandLine, bool success)
{
return predictor => predictor.OnCommandLineExecuted(client, commandLine, success);
}
}
/// <summary>
@@ -161,7 +210,7 @@ namespace System.Management.Automation.Subsystem
/// When the value is greater than 0, it's the number of displayed suggestions from the list returned in <paramref name="session"/>, starting from the index 0.
/// When the value is less than or equal to 0, it means a single suggestion from the list got displayed, and the index is the absolute value.
/// </param>
public static void OnSuggestionDisplayed(string client, Guid predictorId, uint session, int countOrIndex)
public static void OnSuggestionDisplayed(PredictionClient client, Guid predictorId, uint session, int countOrIndex)
{
var predictors = SubsystemManager.GetSubsystems<ICommandPredictor>();
if (predictors.Count == 0)
@@ -171,14 +220,24 @@ namespace System.Management.Automation.Subsystem
foreach (ICommandPredictor predictor in predictors)
{
if (predictor.AcceptFeedback && predictor.Id == predictorId)
if (predictor.Id == predictorId)
{
ThreadPool.QueueUserWorkItem<ICommandPredictor>(
state => state.OnSuggestionDisplayed(client, session, countOrIndex),
predictor,
preferLocal: false);
if (predictor.CanAcceptFeedback(client, PredictorFeedbackKind.SuggestionDisplayed))
{
Action<ICommandPredictor> callBack = GetCallBack(client, session, countOrIndex);
ThreadPool.QueueUserWorkItem<ICommandPredictor>(callBack, predictor, preferLocal: false);
}
break;
}
}
// A local helper function to avoid creating an instance of the generated delegate helper class
// when no predictor is registered, or no registered predictor accepts this feedback.
static Action<ICommandPredictor> GetCallBack(PredictionClient client, uint session, int countOrIndex)
{
return predictor => predictor.OnSuggestionDisplayed(client, session, countOrIndex);
}
}
/// <summary>
@@ -188,7 +247,7 @@ namespace System.Management.Automation.Subsystem
/// <param name="predictorId">The identifier of the predictor whose prediction result was accepted.</param>
/// <param name="session">The mini-session where the accepted suggestion came from.</param>
/// <param name="suggestionText">The accepted suggestion text.</param>
public static void OnSuggestionAccepted(string client, Guid predictorId, uint session, string suggestionText)
public static void OnSuggestionAccepted(PredictionClient client, Guid predictorId, uint session, string suggestionText)
{
Requires.NotNullOrEmpty(suggestionText, nameof(suggestionText));
@@ -200,14 +259,24 @@ namespace System.Management.Automation.Subsystem
foreach (ICommandPredictor predictor in predictors)
{
if (predictor.AcceptFeedback && predictor.Id == predictorId)
if (predictor.Id == predictorId)
{
ThreadPool.QueueUserWorkItem<ICommandPredictor>(
state => state.OnSuggestionAccepted(client, session, suggestionText),
predictor,
preferLocal: false);
if (predictor.CanAcceptFeedback(client, PredictorFeedbackKind.SuggestionAccepted))
{
Action<ICommandPredictor> callBack = GetCallBack(client, session, suggestionText);
ThreadPool.QueueUserWorkItem<ICommandPredictor>(callBack, predictor, preferLocal: false);
}
break;
}
}
// A local helper function to avoid creating an instance of the generated delegate helper class
// when no predictor is registered, or no registered predictor accepts this feedback.
static Action<ICommandPredictor> GetCallBack(PredictionClient client, uint session, string suggestionText)
{
return predictor => predictor.OnSuggestionAccepted(client, session, suggestionText);
}
}
}
}
@@ -9,7 +9,7 @@ using System.Management.Automation.Internal;
using System.Management.Automation.Language;
using System.Threading;
namespace System.Management.Automation.Subsystem
namespace System.Management.Automation.Subsystem.Prediction
{
/// <summary>
/// Interface for implementing a predictor plugin.
@@ -26,57 +26,132 @@ namespace System.Management.Automation.Subsystem
/// </summary>
SubsystemKind ISubsystem.Kind => SubsystemKind.CommandPredictor;
/// <summary>
/// Gets a value indicating whether the predictor supports early processing.
/// </summary>
bool SupportEarlyProcessing { get; }
/// <summary>
/// Gets a value indicating whether the predictor accepts feedback about the previous suggestion.
/// </summary>
bool AcceptFeedback { get; }
/// <summary>
/// A command line was accepted to execute.
/// The predictor can start processing early as needed with the latest history.
/// </summary>
/// <param name="clientId">Represents the client that initiates the call.</param>
/// <param name="history">History command lines provided as references for prediction.</param>
void StartEarlyProcessing(string clientId, IReadOnlyList<string> history);
/// <summary>
/// Get the predictive suggestions. It indicates the start of a suggestion rendering session.
/// </summary>
/// <param name="clientId">Represents the client that initiates the call.</param>
/// <param name="client">Represents the client that initiates the call.</param>
/// <param name="context">The <see cref="PredictionContext"/> object to be used for prediction.</param>
/// <param name="cancellationToken">The cancellation token to cancel the prediction.</param>
/// <returns>An instance of <see cref="SuggestionPackage"/>.</returns>
SuggestionPackage GetSuggestion(string clientId, PredictionContext context, CancellationToken cancellationToken);
SuggestionPackage GetSuggestion(PredictionClient client, PredictionContext context, CancellationToken cancellationToken);
/// <summary>
/// Gets a value indicating whether the predictor accepts a specific kind of feedback.
/// </summary>
/// <param name="client">Represents the client that initiates the call.</param>
/// <param name="feedback">A specific type of feedback.</param>
/// <returns>True or false, to indicate whether the specific feedback is accepted.</returns>
bool CanAcceptFeedback(PredictionClient client, PredictorFeedbackKind feedback);
/// <summary>
/// One or more suggestions provided by the predictor were displayed to the user.
/// </summary>
/// <param name="clientId">Represents the client that initiates the call.</param>
/// <param name="client">Represents the client that initiates the call.</param>
/// <param name="session">The mini-session where the displayed suggestions came from.</param>
/// <param name="countOrIndex">
/// When the value is greater than 0, it's the number of displayed suggestions from the list returned in <paramref name="session"/>, starting from the index 0.
/// When the value is less than or equal to 0, it means a single suggestion from the list got displayed, and the index is the absolute value.
/// </param>
void OnSuggestionDisplayed(string clientId, uint session, int countOrIndex);
void OnSuggestionDisplayed(PredictionClient client, uint session, int countOrIndex);
/// <summary>
/// The suggestion provided by the predictor was accepted.
/// </summary>
/// <param name="clientId">Represents the client that initiates the call.</param>
/// <param name="client">Represents the client that initiates the call.</param>
/// <param name="session">Represents the mini-session where the accepted suggestion came from.</param>
/// <param name="acceptedSuggestion">The accepted suggestion text.</param>
void OnSuggestionAccepted(string clientId, uint session, string acceptedSuggestion);
void OnSuggestionAccepted(PredictionClient client, uint session, string acceptedSuggestion);
/// <summary>
/// A command line was accepted to execute.
/// The predictor can start processing early as needed with the latest history.
/// </summary>
/// <param name="client">Represents the client that initiates the call.</param>
/// <param name="history">History command lines provided as references for prediction.</param>
void OnCommandLineAccepted(PredictionClient client, IReadOnlyList<string> history);
/// <summary>
/// A command line was done execution.
/// </summary>
/// <param name="client">Represents the client that initiates the call.</param>
/// <param name="commandLine">The last accepted command line.</param>
/// <param name="success">Shows whether the execution was successful.</param>
void OnCommandLineExecuted(PredictionClient client, string commandLine, bool success);
}
/// <summary>
/// Kinds of feedback a predictor can choose to accept.
/// </summary>
public enum PredictorFeedbackKind
{
/// <summary>
/// Feedback when one or more suggestions are displayed to the user.
/// </summary>
SuggestionDisplayed,
/// <summary>
/// Feedback when a suggestion is accepted by the user.
/// </summary>
SuggestionAccepted,
/// <summary>
/// Feedback when a command line is accepted by the user.
/// </summary>
CommandLineAccepted,
/// <summary>
/// Feedback when the accepted command line finishes its execution.
/// </summary>
CommandLineExecuted,
}
/// <summary>
/// Kinds of prediction clients.
/// </summary>
public enum PredictionClientKind
{
/// <summary>
/// A terminal client, representing the command-line experience.
/// </summary>
Terminal,
/// <summary>
/// An editor client, representing the editor experience.
/// </summary>
Editor,
}
/// <summary>
/// The class represents a client that interacts with predictors.
/// </summary>
public sealed class PredictionClient
{
/// <summary>
/// Gets the client name.
/// </summary>
public string Name { get; }
/// <summary>
/// Gets the client kind.
/// </summary>
public PredictionClientKind Kind { get; }
/// <summary>
/// Initializes a new instance of the <see cref="PredictionClient"/> class.
/// </summary>
/// <param name="name">Name of the interactive client.</param>
/// <param name="kind">Kind of the interactive client.</param>
public PredictionClient(string name, PredictionClientKind kind)
{
Name = name;
Kind = kind;
}
}
/// <summary>
/// Context information about the user input.
/// </summary>
public class PredictionContext
public sealed class PredictionContext
{
/// <summary>
/// Gets the abstract syntax tree (AST) generated from parsing the user input.
@@ -8,6 +8,7 @@ using System.Collections.Generic;
using System.Collections.ObjectModel;
using System.Linq;
using System.Management.Automation.Internal;
using System.Management.Automation.Subsystem.Prediction;
namespace System.Management.Automation.Subsystem
{
+67 -35
View File
@@ -5,6 +5,7 @@ using System;
using System.Collections.Generic;
using System.Management.Automation.Language;
using System.Management.Automation.Subsystem;
using System.Management.Automation.Subsystem.Prediction;
using System.Threading;
using Xunit;
@@ -18,6 +19,8 @@ namespace PSTests.Sequential
public List<string> History { get; }
public List<string> Results { get; }
public List<string> AcceptedSuggestions { get; }
public List<string> DisplayedSuggestions { get; }
@@ -47,39 +50,30 @@ namespace PSTests.Sequential
_delay = delay;
History = new List<string>();
Results = new List<string>();
AcceptedSuggestions = new List<string>();
DisplayedSuggestions = new List<string>();
}
public void Clear()
{
History.Clear();
Results.Clear();
AcceptedSuggestions.Clear();
DisplayedSuggestions.Clear();
}
#region "Interface implementation"
public Guid Id => _id;
public string Name => _name;
public string Description => _description;
bool ICommandPredictor.SupportEarlyProcessing => true;
bool ICommandPredictor.CanAcceptFeedback(PredictionClient client, PredictorFeedbackKind feedback) => true;
bool ICommandPredictor.AcceptFeedback => true;
public void StartEarlyProcessing(string clientId, IReadOnlyList<string> history)
{
foreach (string item in history)
{
History.Add($"{clientId}-{item}");
}
}
public void OnSuggestionDisplayed(string clientId, uint session, int countOrIndex)
{
DisplayedSuggestions.Add($"{clientId}-{session}-{countOrIndex}");
}
public void OnSuggestionAccepted(string clientId, uint session, string acceptedSuggestion)
{
AcceptedSuggestions.Add($"{clientId}-{session}-{acceptedSuggestion}");
}
public SuggestionPackage GetSuggestion(string clientId, PredictionContext context, CancellationToken cancellationToken)
public SuggestionPackage GetSuggestion(PredictionClient client, PredictionContext context, CancellationToken cancellationToken)
{
if (_delay)
{
@@ -92,18 +86,44 @@ namespace PSTests.Sequential
var userInput = context.InputAst.Extent.Text;
var entries = new List<PredictiveSuggestion>
{
new PredictiveSuggestion($"'{userInput}' from '{clientId}' - TEST-1 from {Name}"),
new PredictiveSuggestion($"'{userInput}' from '{clientId}' - TeSt-2 from {Name}"),
new PredictiveSuggestion($"'{userInput}' from '{client.Name}' - TEST-1 from {Name}"),
new PredictiveSuggestion($"'{userInput}' from '{client.Name}' - TeSt-2 from {Name}"),
};
return new SuggestionPackage(56, entries);
}
public void OnSuggestionDisplayed(PredictionClient client, uint session, int countOrIndex)
{
DisplayedSuggestions.Add($"{client.Name}-{session}-{countOrIndex}");
}
public void OnSuggestionAccepted(PredictionClient client, uint session, string acceptedSuggestion)
{
AcceptedSuggestions.Add($"{client.Name}-{session}-{acceptedSuggestion}");
}
public void OnCommandLineAccepted(PredictionClient client, IReadOnlyList<string> history)
{
foreach (string item in history)
{
History.Add($"{client.Name}-{item}");
}
}
public void OnCommandLineExecuted(PredictionClient client, string commandLine, bool success)
{
Results.Add($"{client.Name}-{commandLine}-{success}");
}
#endregion
}
public static class CommandPredictionTests
{
private const string Client = "PredictionTest";
private const uint Session = 56;
private static PredictionClient predClient = new(Client, PredictionClientKind.Terminal);
[Fact]
public static void PredictInput()
@@ -114,7 +134,7 @@ namespace PSTests.Sequential
Ast ast = Parser.ParseInput(Input, out Token[] tokens, out _);
// Returns null when no predictor implementation registered
List<PredictionResult> results = CommandPrediction.PredictInput(Client, ast, tokens).Result;
List<PredictionResult> results = CommandPrediction.PredictInputAsync(predClient, ast, tokens).Result;
Assert.Null(results);
try
@@ -127,7 +147,7 @@ namespace PSTests.Sequential
// cannot finish before the specified timeout.
// The specified timeout is exaggerated to make the test reliable.
// xUnit must spin up a lot tasks, which makes the test unreliable when the time difference between 'delay' and 'timeout' is small.
results = CommandPrediction.PredictInput(Client, ast, tokens, millisecondsTimeout: 1000).Result;
results = CommandPrediction.PredictInputAsync(predClient, ast, tokens, millisecondsTimeout: 1000).Result;
Assert.Single(results);
PredictionResult res = results[0];
@@ -140,7 +160,7 @@ namespace PSTests.Sequential
// Expect the results from both 'slow' and 'fast' predictors
// Same here -- the specified timeout is exaggerated to make the test reliable.
// xUnit must spin up a lot tasks, which makes the test unreliable when the time difference between 'delay' and 'timeout' is small.
results = CommandPrediction.PredictInput(Client, ast, tokens, millisecondsTimeout: 4000).Result;
results = CommandPrediction.PredictInputAsync(predClient, ast, tokens, millisecondsTimeout: 4000).Result;
Assert.Equal(2, results.Count);
PredictionResult res1 = results[0];
@@ -170,6 +190,9 @@ namespace PSTests.Sequential
MyPredictor slow = MyPredictor.SlowPredictor;
MyPredictor fast = MyPredictor.FastPredictor;
slow.Clear();
fast.Clear();
try
{
// Register 2 predictor implementations
@@ -179,16 +202,19 @@ namespace PSTests.Sequential
var history = new[] { "hello", "world" };
var ids = new HashSet<Guid> { slow.Id, fast.Id };
CommandPrediction.OnCommandLineAccepted(Client, history);
CommandPrediction.OnSuggestionDisplayed(Client, slow.Id, Session, 2);
CommandPrediction.OnSuggestionDisplayed(Client, fast.Id, Session, -1);
CommandPrediction.OnSuggestionAccepted(Client, slow.Id, Session, "Yeah");
CommandPrediction.OnCommandLineAccepted(predClient, history);
CommandPrediction.OnCommandLineExecuted(predClient, "last_input", true);
CommandPrediction.OnSuggestionDisplayed(predClient, slow.Id, Session, 2);
CommandPrediction.OnSuggestionDisplayed(predClient, fast.Id, Session, -1);
CommandPrediction.OnSuggestionAccepted(predClient, slow.Id, Session, "Yeah");
// The calls to 'StartEarlyProcessing' and 'OnSuggestionAccepted' are queued in thread pool,
// so we wait a bit to make sure the calls are done.
while (slow.History.Count == 0 || slow.AcceptedSuggestions.Count == 0)
// The feedback calls are queued in thread pool, so let's wait a bit to make sure the calls are done.
while (slow.History.Count == 0 || fast.History.Count == 0 ||
slow.Results.Count == 0 || fast.Results.Count == 0 ||
slow.DisplayedSuggestions.Count == 0 || fast.DisplayedSuggestions.Count == 0 ||
slow.AcceptedSuggestions.Count == 0)
{
Thread.Sleep(10);
Thread.Sleep(100);
}
Assert.Equal(2, slow.History.Count);
@@ -199,6 +225,12 @@ namespace PSTests.Sequential
Assert.Equal($"{Client}-{history[0]}", fast.History[0]);
Assert.Equal($"{Client}-{history[1]}", fast.History[1]);
Assert.Single(slow.Results);
Assert.Equal($"{Client}-last_input-True", slow.Results[0]);
Assert.Single(fast.Results);
Assert.Equal($"{Client}-last_input-True", fast.Results[0]);
Assert.Single(slow.DisplayedSuggestions);
Assert.Equal($"{Client}-{Session}-2", slow.DisplayedSuggestions[0]);
+3 -1
View File
@@ -4,6 +4,7 @@
using System;
using System.Collections.ObjectModel;
using System.Management.Automation.Subsystem;
using System.Management.Automation.Subsystem.Prediction;
using System.Threading;
using Xunit;
@@ -97,8 +98,9 @@ namespace PSTests.Sequential
const string Client = "SubsystemTest";
const string Input = "Hello world";
var predClient = new PredictionClient(Client, PredictionClientKind.Terminal);
var predCxt = PredictionContext.Create(Input);
var results = impl.GetSuggestion(Client, predCxt, CancellationToken.None);
var results = impl.GetSuggestion(predClient, predCxt, CancellationToken.None);
Assert.Equal($"'{Input}' from '{Client}' - TEST-1 from {impl.Name}", results.SuggestionEntries[0].SuggestionText);
Assert.Equal($"'{Input}' from '{Client}' - TeSt-2 from {impl.Name}", results.SuggestionEntries[1].SuggestionText);