diff --git a/README.md b/README.md index c0894821..9d80bb5b 100644 --- a/README.md +++ b/README.md @@ -161,7 +161,57 @@ foreach (AiClient::defaultRegistry()->findModelsMetadataForSupport($requirements } ``` -See the [`PromptBuilder` class](https://github.com/WordPress/php-ai-client/blob/trunk/src/Builders/PromptBuilder.php) and the [`EmbeddingBuilder` class](https://github.com/WordPress/php-ai-client/blob/trunk/src/Builders/EmbeddingBuilder.php) and their public methods for all the ways you can configure generation. +### Classification using any compatible model + +Classification models answer typed questions about the state you give them rather than generating content. There are three types of question: + +* A `binary` question asks whether a statement is true, and is answered with the probability that it is. +* A `choice` question lists named options, each with a description of when it applies, and is answered with the key of the chosen option. +* A `score` question lists the levels of a scale from lowest to highest, and is answered with a position on that scale, from `0` for the lowest level to one less than the number of levels. The position may fall between two levels. + +Answers to `choice` and `score` questions may also carry the model's confidence and a probability for each option or level, depending on what the model reports. Every answer is checked against its question, so an answer of the wrong type, an option the question does not have, or a position outside the scale raises a `RuntimeException`. + +As with content generation, a suitable model is discovered if you do not name one. Classification requires a registered provider with at least one model that supports it. If there is none, `classifyResult()` throws an `InvalidArgumentException`, and `isSupported()` returns `false`. + +```php +use WordPress\AiClient\AiClient; +use WordPress\AiClient\Providers\Models\Classification\DTO\ClassificationQuestion; +use WordPress\AiClient\Providers\Models\Classification\Enums\ClassificationQuestionTypeEnum; + +$result = AiClient::classify(['comment' => $commentText]) + ->withQuestion( + 'spam', + new ClassificationQuestion(ClassificationQuestionTypeEnum::binary(), 'Is this comment spam?') + ) + ->withQuestion( + 'route', + new ClassificationQuestion( + ClassificationQuestionTypeEnum::choice(), + 'How should a moderator handle it?', + [ + 'approve' => 'The comment is fine to publish.', + 'hold' => 'The comment needs a closer look.', + 'trash' => 'The comment should be removed.', + ] + ) + ) + ->withQuestion( + 'tone', + new ClassificationQuestion( + ClassificationQuestionTypeEnum::score(), + 'How civil is the comment?', + ['Hostile.', 'Neutral.', 'Friendly.'] + ) + ) + ->classifyResult(); + +$spamProbability = $result->getAnswer('spam')->getProbability(); // e.g. 0.03 +$route = $result->getAnswer('route')->getChoice(); // e.g. 'approve' +$confidence = $result->getAnswer('route')->getConfidence(); // e.g. 0.91, or null if not reported +$tone = $result->getAnswer('tone')->getScore(); // e.g. 1.7, between 'Neutral.' and 'Friendly.' +``` + +See the [`PromptBuilder` class](https://github.com/WordPress/php-ai-client/blob/trunk/src/Builders/PromptBuilder.php), the [`EmbeddingBuilder` class](https://github.com/WordPress/php-ai-client/blob/trunk/src/Builders/EmbeddingBuilder.php), and the [`ClassificationBuilder` class](https://github.com/WordPress/php-ai-client/blob/trunk/src/Builders/ClassificationBuilder.php) and their public methods for all the ways you can configure generation. **More documentation is coming soon.** @@ -175,6 +225,8 @@ The AI Client supports PSR-14 event dispatching for prompt lifecycle events. Thi - `AfterGenerateResultEvent` - Dispatched after a result is received from the model - `BeforeGenerateEmbeddingEvent` - Dispatched before embedding inputs are sent to the model - `AfterGenerateEmbeddingEvent` - Dispatched after an embedding result is received from the model +- `BeforeClassifyEvent` - Dispatched before classification state and questions are sent to the model +- `AfterClassifyEvent` - Dispatched after the model's answers have been received and checked against the questions **Important:** Event listeners should not return a value, as they will be ignored. In order to modify data that is passed with the event object, you need to rely on setters on the event object. Any event data for which there are no setters on the event object is meant to be immutable or, in other words, read-only for the event listener. diff --git a/docs/ARCHITECTURE.md b/docs/ARCHITECTURE.md index f6635276..766d8b1d 100644 --- a/docs/ARCHITECTURE.md +++ b/docs/ARCHITECTURE.md @@ -22,6 +22,8 @@ The two builders differ fundamentally in how they arrive at a model: Both builders accumulate model configuration through the shared `ModelConfigurationTrait`, which provides `usingModelConfig()`. +Classification uses a third builder, `ClassificationBuilder`, because classification models take state and typed questions rather than a prompt, and answer each question with a probability, an option, or a position on a scale instead of generating content. Like `PromptBuilder`, it **resolves** a model through the `ModelResolutionTrait`, since answers are not tied to the model that produced them. + ### Code examples The following examples indicate how this SDK could eventually be used. diff --git a/src/AiClient.php b/src/AiClient.php index 23b308a9..b3c410b9 100644 --- a/src/AiClient.php +++ b/src/AiClient.php @@ -6,6 +6,7 @@ use Psr\EventDispatcher\EventDispatcherInterface; use Psr\SimpleCache\CacheInterface; +use WordPress\AiClient\Builders\ClassificationBuilder; use WordPress\AiClient\Builders\EmbeddingBuilder; use WordPress\AiClient\Builders\PromptBuilder; use WordPress\AiClient\Common\Exception\InvalidArgumentException; @@ -271,6 +272,28 @@ public static function input($input = null, ?ProviderRegistry $registry = null): ); } + /** + * Creates a classification builder for fluent classification. + * + * Classification answers typed questions about the given state rather than generating content. Add + * questions via withQuestion() on the returned builder, then chain classifyResult() to get one answer + * per question. + * + * @since n.e.x.t + * + * @param array|null $state Optional initial state to classify. + * @param ProviderRegistry|null $registry Optional custom registry. If null, uses default. + * @return ClassificationBuilder The classification builder instance. + */ + public static function classify(?array $state = null, ?ProviderRegistry $registry = null): ClassificationBuilder + { + return new ClassificationBuilder( + $registry ?? self::defaultRegistry(), + $state, + self::$eventDispatcher + ); + } + /** * Generates content using a unified API that automatically detects model capabilities. * diff --git a/src/Builders/ClassificationBuilder.php b/src/Builders/ClassificationBuilder.php new file mode 100644 index 00000000..cfa735d0 --- /dev/null +++ b/src/Builders/ClassificationBuilder.php @@ -0,0 +1,327 @@ + The state to classify. + */ + protected array $state = []; + + /** + * @var array The questions to answer, keyed by question key. + */ + protected array $questions = []; + + /** + * @var EventDispatcherInterface|null The event dispatcher for classification lifecycle events. + */ + private ?EventDispatcherInterface $eventDispatcher; + + /** + * Constructor. + * + * @since n.e.x.t + * + * @param ProviderRegistry $registry The provider registry for finding suitable models. + * @param array|null $state Optional initial state to classify. + * @param EventDispatcherInterface|null $eventDispatcher Optional event dispatcher for lifecycle events. + */ + public function __construct( + ProviderRegistry $registry, + ?array $state = null, + ?EventDispatcherInterface $eventDispatcher = null + ) { + $this->modelConfig = new ModelConfig(); + $this->modelResolver = new ModelResolver($registry); + $this->eventDispatcher = $eventDispatcher; + + if ($state !== null) { + $this->withState($state); + } + } + + /** + * Creates a deep clone of this builder. + * + * Questions are immutable and therefore shared between clones. The event dispatcher is a service object + * and is intentionally NOT cloned. + * + * @since n.e.x.t + */ + public function __clone() + { + $this->modelConfig = clone $this->modelConfig; + $this->modelResolver = clone $this->modelResolver; + } + + /** + * Adds state to classify. + * + * Keys that are already present are overwritten. + * + * @since n.e.x.t + * + * @param array $state The state to add, keyed by name. + * @return self + * @throws InvalidArgumentException If the state is empty or not keyed by name. + */ + public function withState(array $state): self + { + if ($state === []) { + throw new InvalidArgumentException('Classification state cannot be empty.'); + } + + foreach (array_keys($state) as $key) { + // PHP converts integer-like string keys to integers, so those are rejected along with lists. + if (!is_string($key) || trim($key) === '') { + throw new InvalidArgumentException('Classification state keys must be non-empty, non-integer strings.'); + } + } + + $this->state = array_replace($this->state, $state); + + return $this; + } + + /** + * Adds a question to answer about the state. + * + * A question added with an existing key replaces the previous question. + * + * @since n.e.x.t + * + * @param string $key The key identifying the question and its answer. + * @param ClassificationQuestion $question The question. + * @return self + * @throws InvalidArgumentException If the key is empty or an integer. + */ + public function withQuestion(string $key, ClassificationQuestion $question): self + { + // PHP converts integer-like array keys such as "1" to integers, so the questions could no longer be + // keyed by name. + if (trim($key) === '' || (string) (int) $key === $key) { + throw new InvalidArgumentException('Classification question keys must be non-empty, non-integer strings.'); + } + + $this->questions[$key] = $question; + + return $this; + } + + /** + * Checks whether a model is available for classification with the current configuration. + * + * @since n.e.x.t + * + * @return bool True if the set model or any registered model supports classification. + */ + public function isSupported(): bool + { + return $this->modelResolver->isSupported($this->getModelRequirements()); + } + + /** + * Classifies the state by answering the configured questions. + * + * @since n.e.x.t + * + * @return ClassificationResult The result containing one answer per question. + * @throws InvalidArgumentException If no state or questions are configured, or no suitable model is found. + * @throws RuntimeException If the model does not support classification, or does not give a valid answer + * to every question. + */ + public function classifyResult(): ClassificationResult + { + if ($this->state === []) { + throw new InvalidArgumentException('Cannot classify empty state. Add state using withState().'); + } + + if ($this->questions === []) { + throw new InvalidArgumentException( + 'Cannot classify without questions. Add questions using withQuestion().' + ); + } + + $capability = CapabilityEnum::classification(); + $model = $this->modelResolver->resolve($this->getModelRequirements(), $this->modelConfig); + + if (!$model instanceof ClassificationModelInterface) { + throw new RuntimeException( + sprintf( + 'Model "%s" does not support classification.', + $model->metadata()->getId() + ) + ); + } + + $this->dispatchEvent(new BeforeClassifyEvent($this->state, $this->questions, $model, $capability)); + + $result = $model->classifyResult($this->state, $this->questions); + + $this->assertValidAnswers($result); + + $this->dispatchEvent( + new AfterClassifyEvent($this->state, $this->questions, $model, $capability, $result) + ); + + return $result; + } + + /** + * Gets the requirements a model must meet to classify with the current configuration. + * + * @since n.e.x.t + * + * @return ModelRequirements The model requirements. + */ + private function getModelRequirements(): ModelRequirements + { + // Classification takes state and questions rather than messages, so only the capability and the + // model configuration contribute requirements. + return ModelRequirements::fromPromptData(CapabilityEnum::classification(), [], $this->modelConfig); + } + + /** + * Asserts that the result holds a valid answer to every question. + * + * @since n.e.x.t + * + * @param ClassificationResult $result The result from the model. + * @return void + * @throws RuntimeException If a question is unanswered, or an answer does not fit its question. + */ + private function assertValidAnswers(ClassificationResult $result): void + { + $answers = $result->getAnswers(); + + // Answers map to questions by key, so every question must be answered. + $unanswered = array_diff_key($this->questions, $answers); + if ($unanswered !== []) { + throw new RuntimeException( + sprintf( + 'The model did not answer the following questions: %s.', + implode(', ', array_keys($unanswered)) + ) + ); + } + + foreach ($this->questions as $key => $question) { + $problem = $this->describeInvalidAnswer($question, $answers[$key]); + if ($problem !== null) { + throw new RuntimeException( + sprintf('The model gave an invalid answer to the question "%s". %s', $key, $problem) + ); + } + } + } + + /** + * Describes why an answer does not fit the question it answers. + * + * @since n.e.x.t + * + * @param ClassificationQuestion $question The question. + * @param ClassificationAnswer $answer The answer to the question. + * @return string|null A description of the problem, or null if the answer fits the question. + */ + private function describeInvalidAnswer(ClassificationQuestion $question, ClassificationAnswer $answer): ?string + { + $type = $question->getType(); + + if (!$answer->getType()->is($type)) { + return sprintf( + 'Expected an answer to a %s question, but received an answer to a %s question.', + $type->value, + $answer->getType()->value + ); + } + + $criteria = $question->getCriteria(); + $probabilities = $answer->getProbabilities(); + + if ($type->isChoice()) { + if (!array_key_exists($answer->getChoice(), $criteria)) { + return sprintf('"%s" is not one of its options.', $answer->getChoice()); + } + + $unknownOptions = array_diff_key($probabilities, $criteria); + if ($unknownOptions !== []) { + return sprintf( + 'It has probabilities for options the question does not have: %s.', + implode(', ', array_keys($unknownOptions)) + ); + } + } + + if ($type->isScore()) { + $levelCount = count($criteria); + + if ($answer->getScore() > $levelCount - 1) { + return sprintf( + 'The score %s is outside the scale, which runs from 0 to %d.', + $answer->getScore(), + $levelCount - 1 + ); + } + + if ($probabilities !== [] && count($probabilities) !== $levelCount) { + return sprintf( + 'It has probabilities for %d levels, but the question has %d.', + count($probabilities), + $levelCount + ); + } + } + + return null; + } + + /** + * Dispatches an event if an event dispatcher is registered. + * + * @since n.e.x.t + * + * @param object $event The event to dispatch. + * @return void + */ + private function dispatchEvent(object $event): void + { + if ($this->eventDispatcher !== null) { + $this->eventDispatcher->dispatch($event); + } + } +} diff --git a/src/Events/AfterClassifyEvent.php b/src/Events/AfterClassifyEvent.php new file mode 100644 index 00000000..9779b13d --- /dev/null +++ b/src/Events/AfterClassifyEvent.php @@ -0,0 +1,144 @@ + The state that was sent to the model. + */ + private array $state; + + /** + * @var array The questions that were sent to the model, keyed by question key. + */ + private array $questions; + + /** + * @var ModelInterface The model that classified the state. + */ + private ModelInterface $model; + + /** + * @var CapabilityEnum The capability that was used for classification. + */ + private CapabilityEnum $capability; + + /** + * @var ClassificationResult The result from the model. + */ + private ClassificationResult $result; + + /** + * Constructor. + * + * @since n.e.x.t + * + * @param array $state The state that was sent to the model. + * @param array $questions The questions that were sent to the model, keyed + * by question key. + * @param ModelInterface $model The model that classified the state. + * @param CapabilityEnum $capability The capability that was used for classification. + * @param ClassificationResult $result The result from the model. + */ + public function __construct( + array $state, + array $questions, + ModelInterface $model, + CapabilityEnum $capability, + ClassificationResult $result + ) { + $this->state = $state; + $this->questions = $questions; + $this->model = $model; + $this->capability = $capability; + $this->result = $result; + } + + /** + * Gets the state that was sent to the model. + * + * @since n.e.x.t + * + * @return array The state. + */ + public function getState(): array + { + return $this->state; + } + + /** + * Gets the questions that were sent to the model. + * + * @since n.e.x.t + * + * @return array The questions, keyed by question key. + */ + public function getQuestions(): array + { + return $this->questions; + } + + /** + * Gets the model that classified the state. + * + * @since n.e.x.t + * + * @return ModelInterface The model. + */ + public function getModel(): ModelInterface + { + return $this->model; + } + + /** + * Gets the capability that was used for classification. + * + * @since n.e.x.t + * + * @return CapabilityEnum The capability. + */ + public function getCapability(): CapabilityEnum + { + return $this->capability; + } + + /** + * Gets the result from the model. + * + * @since n.e.x.t + * + * @return ClassificationResult The result. + */ + public function getResult(): ClassificationResult + { + return $this->result; + } + + /** + * Performs a deep clone of the event. + * + * @since n.e.x.t + */ + public function __clone() + { + $clonedQuestions = []; + foreach ($this->questions as $key => $question) { + $clonedQuestions[$key] = clone $question; + } + $this->questions = $clonedQuestions; + $this->result = clone $this->result; + } +} diff --git a/src/Events/BeforeClassifyEvent.php b/src/Events/BeforeClassifyEvent.php new file mode 100644 index 00000000..c448c908 --- /dev/null +++ b/src/Events/BeforeClassifyEvent.php @@ -0,0 +1,118 @@ + The state to be sent to the model. + */ + private array $state; + + /** + * @var array The questions to be sent to the model, keyed by question key. + */ + private array $questions; + + /** + * @var ModelInterface The model that will classify the state. + */ + private ModelInterface $model; + + /** + * @var CapabilityEnum The capability being used for classification. + */ + private CapabilityEnum $capability; + + /** + * Constructor. + * + * @since n.e.x.t + * + * @param array $state The state to be sent to the model. + * @param array $questions The questions to be sent to the model, keyed by + * question key. + * @param ModelInterface $model The model that will classify the state. + * @param CapabilityEnum $capability The capability being used for classification. + */ + public function __construct(array $state, array $questions, ModelInterface $model, CapabilityEnum $capability) + { + $this->state = $state; + $this->questions = $questions; + $this->model = $model; + $this->capability = $capability; + } + + /** + * Gets the state to be sent to the model. + * + * @since n.e.x.t + * + * @return array The state. + */ + public function getState(): array + { + return $this->state; + } + + /** + * Gets the questions to be sent to the model. + * + * @since n.e.x.t + * + * @return array The questions, keyed by question key. + */ + public function getQuestions(): array + { + return $this->questions; + } + + /** + * Gets the model that will classify the state. + * + * @since n.e.x.t + * + * @return ModelInterface The model. + */ + public function getModel(): ModelInterface + { + return $this->model; + } + + /** + * Gets the capability being used for classification. + * + * @since n.e.x.t + * + * @return CapabilityEnum The capability. + */ + public function getCapability(): CapabilityEnum + { + return $this->capability; + } + + /** + * Performs a deep clone of the event. + * + * @since n.e.x.t + */ + public function __clone() + { + $clonedQuestions = []; + foreach ($this->questions as $key => $question) { + $clonedQuestions[$key] = clone $question; + } + $this->questions = $clonedQuestions; + } +} diff --git a/src/Providers/Models/Classification/Contracts/ClassificationModelInterface.php b/src/Providers/Models/Classification/Contracts/ClassificationModelInterface.php new file mode 100644 index 00000000..e457c123 --- /dev/null +++ b/src/Providers/Models/Classification/Contracts/ClassificationModelInterface.php @@ -0,0 +1,30 @@ + $state The state to classify, keyed by name. + * @param array $questions The questions to answer, keyed by question key. + * @return ClassificationResult Result containing one answer per question, keyed by question key. + */ + public function classifyResult(array $state, array $questions): ClassificationResult; +} diff --git a/src/Providers/Models/Classification/DTO/ClassificationQuestion.php b/src/Providers/Models/Classification/DTO/ClassificationQuestion.php new file mode 100644 index 00000000..f9da7987 --- /dev/null +++ b/src/Providers/Models/Classification/DTO/ClassificationQuestion.php @@ -0,0 +1,230 @@ +|list + * } + * + * @extends AbstractDataTransferObject + */ +class ClassificationQuestion extends AbstractDataTransferObject +{ + public const KEY_TYPE = 'type'; + public const KEY_INSTRUCTIONS = 'instructions'; + public const KEY_CRITERIA = 'criteria'; + + /** + * @var ClassificationQuestionTypeEnum The question type. + */ + private ClassificationQuestionTypeEnum $type; + + /** + * @var string The question to answer. + */ + private string $instructions; + + /** + * @var array|list The option descriptions keyed by option key, or the level + * descriptions from lowest to highest. + */ + private array $criteria; + + /** + * Constructor. + * + * @since n.e.x.t + * + * @param ClassificationQuestionTypeEnum $type The question type. + * @param string $instructions The question to answer. + * @param array|list $criteria For a choice question, the option descriptions keyed + * by option key. For a score question, the level + * descriptions from lowest to highest. Empty for a + * binary question. + * + * @throws InvalidArgumentException If the instructions are empty or the criteria do not suit the type. + */ + public function __construct(ClassificationQuestionTypeEnum $type, string $instructions, array $criteria = []) + { + if (trim($instructions) === '') { + throw new InvalidArgumentException('Classification question instructions cannot be empty.'); + } + + if ($type->isBinary()) { + if ($criteria !== []) { + throw new InvalidArgumentException('Binary classification questions do not accept criteria.'); + } + } elseif (count($criteria) < 2) { + throw new InvalidArgumentException( + sprintf('Classification questions of type "%s" require at least two criteria.', $type->value) + ); + } + + if ($type->isChoice()) { + foreach (array_keys($criteria) as $optionKey) { + // PHP converts integer-like string keys to integers, so those are rejected along with lists. + if (!is_string($optionKey) || trim($optionKey) === '') { + throw new InvalidArgumentException( + 'Choice classification question option keys must be non-empty, non-integer strings.' + ); + } + } + } + + if ($type->isScore() && !array_is_list($criteria)) { + throw new InvalidArgumentException( + 'Score classification question criteria must be a list of level descriptions, ' + . 'from lowest to highest.' + ); + } + + foreach ($criteria as $description) { + if (!is_string($description) || trim($description) === '') { + throw new InvalidArgumentException( + 'Classification question criteria descriptions must be non-empty strings.' + ); + } + } + + $this->type = $type; + $this->instructions = $instructions; + $this->criteria = $criteria; + } + + /** + * Gets the question type. + * + * @since n.e.x.t + * + * @return ClassificationQuestionTypeEnum The question type. + */ + public function getType(): ClassificationQuestionTypeEnum + { + return $this->type; + } + + /** + * Gets the question to answer. + * + * @since n.e.x.t + * + * @return string The question instructions. + */ + public function getInstructions(): string + { + return $this->instructions; + } + + /** + * Gets the criteria to answer the question with. + * + * @since n.e.x.t + * + * @return array|list The option descriptions keyed by option key for a choice + * question, the level descriptions from lowest to highest for + * a score question, or an empty array for a binary question. + */ + public function getCriteria(): array + { + return $this->criteria; + } + + /** + * {@inheritDoc} + * + * @since n.e.x.t + */ + public static function getJsonSchema(): array + { + return [ + 'type' => 'object', + 'properties' => [ + self::KEY_TYPE => [ + 'type' => 'string', + 'enum' => ClassificationQuestionTypeEnum::getValues(), + 'description' => 'The question type.', + ], + self::KEY_INSTRUCTIONS => [ + 'type' => 'string', + 'description' => 'The question to answer.', + ], + self::KEY_CRITERIA => [ + 'oneOf' => [ + [ + 'type' => 'object', + 'additionalProperties' => [ + 'type' => 'string', + ], + 'minProperties' => 2, + ], + [ + 'type' => 'array', + 'items' => [ + 'type' => 'string', + ], + 'minItems' => 2, + ], + ], + 'description' => 'The option descriptions keyed by option key for a choice question, or the ' + . 'level descriptions from lowest to highest for a score question.', + ], + ], + 'required' => [self::KEY_TYPE, self::KEY_INSTRUCTIONS], + ]; + } + + /** + * {@inheritDoc} + * + * @since n.e.x.t + * + * @return ClassificationQuestionArrayShape + */ + public function toArray(): array + { + $data = [ + self::KEY_TYPE => $this->type->value, + self::KEY_INSTRUCTIONS => $this->instructions, + ]; + + if ($this->criteria !== []) { + $data[self::KEY_CRITERIA] = $this->criteria; + } + + return $data; + } + + /** + * {@inheritDoc} + * + * @since n.e.x.t + */ + public static function fromArray(array $array): self + { + static::validateFromArrayData($array, [self::KEY_TYPE, self::KEY_INSTRUCTIONS]); + + return new self( + ClassificationQuestionTypeEnum::from($array[self::KEY_TYPE]), + $array[self::KEY_INSTRUCTIONS], + $array[self::KEY_CRITERIA] ?? [] + ); + } +} diff --git a/src/Providers/Models/Classification/Enums/ClassificationQuestionTypeEnum.php b/src/Providers/Models/Classification/Enums/ClassificationQuestionTypeEnum.php new file mode 100644 index 00000000..56564426 --- /dev/null +++ b/src/Providers/Models/Classification/Enums/ClassificationQuestionTypeEnum.php @@ -0,0 +1,37 @@ +|list + * } + * + * @extends AbstractDataTransferObject + */ +class ClassificationAnswer extends AbstractDataTransferObject +{ + public const KEY_TYPE = 'type'; + public const KEY_VALUE = 'value'; + public const KEY_CONFIDENCE = 'confidence'; + public const KEY_PROBABILITIES = 'probabilities'; + + /** + * @var ClassificationQuestionTypeEnum The type of the question answered. + */ + private ClassificationQuestionTypeEnum $type; + + /** + * @var float|string The probability, the chosen option key, or the position on the scale. + */ + private $value; + + /** + * @var float|null The model's confidence in the answer, if reported. + */ + private ?float $confidence; + + /** + * @var array|list The probability of each option or level, if reported. + */ + private array $probabilities; + + /** + * Constructor. + * + * @since n.e.x.t + * + * @param ClassificationQuestionTypeEnum $type The type of the question answered. + * @param float|string $value The probability for a binary question, the chosen option key for a choice + * question, or the position on the scale for a score question. + * @param float|null $confidence The model's confidence in the answer, if reported. + * @param array|list $probabilities For a choice question, the probability of each + * option keyed by option key. For a score question, + * the probability of each level from lowest to + * highest. Empty if not reported, and always empty + * for a binary question. + * + * @throws InvalidArgumentException If the value, confidence, or probabilities do not suit the type, or a + * confidence or probability is not between 0 and 1. + */ + public function __construct( + ClassificationQuestionTypeEnum $type, + $value, + ?float $confidence = null, + array $probabilities = [] + ) { + if ($type->isBinary()) { + $value = self::validateProbability($value, 'Classification answer probability'); + } elseif ($type->isChoice()) { + if (!is_string($value) || trim($value) === '') { + throw new InvalidArgumentException('Classification answer choice must be a non-empty option key.'); + } + } else { + if ((!is_int($value) && !is_float($value)) || !is_finite((float) $value) || $value < 0) { + throw new InvalidArgumentException('Classification answer score must be a non-negative number.'); + } + + $value = (float) $value; + } + + if ($confidence !== null) { + $confidence = self::validateProbability($confidence, 'Classification answer confidence'); + } + + if ($probabilities !== []) { + if ($type->isBinary()) { + throw new InvalidArgumentException( + 'Classification answers to binary questions do not accept probabilities, ' + . 'as their value is the probability.' + ); + } + + if ($type->isScore() && !array_is_list($probabilities)) { + throw new InvalidArgumentException( + 'Classification answer probabilities for a score question must be a list, ' + . 'from the lowest level to the highest.' + ); + } + + foreach ($probabilities as $key => $probability) { + if ($type->isChoice() && (!is_string($key) || trim($key) === '')) { + throw new InvalidArgumentException( + 'Classification answer probabilities for a choice question must be keyed by option key.' + ); + } + + $probabilities[$key] = self::validateProbability( + $probability, + 'Classification answer probabilities' + ); + } + } + + $this->type = $type; + $this->value = $value; + $this->confidence = $confidence; + $this->probabilities = $probabilities; + } + + /** + * Gets the type of the question answered. + * + * @since n.e.x.t + * + * @return ClassificationQuestionTypeEnum The question type. + */ + public function getType(): ClassificationQuestionTypeEnum + { + return $this->type; + } + + /** + * Gets the probability that the statement of a binary question is true. + * + * @since n.e.x.t + * + * @return float The probability, from 0 to 1. + * @throws RuntimeException If this is not an answer to a binary question. + */ + public function getProbability(): float + { + $this->assertType(ClassificationQuestionTypeEnum::binary()); + + return (float) $this->value; + } + + /** + * Gets the chosen option of a choice question. + * + * @since n.e.x.t + * + * @return string The option key. + * @throws RuntimeException If this is not an answer to a choice question. + */ + public function getChoice(): string + { + $this->assertType(ClassificationQuestionTypeEnum::choice()); + + return (string) $this->value; + } + + /** + * Gets the position on the scale of a score question. + * + * @since n.e.x.t + * + * @return float The position, from 0 for the lowest level to one less than the number of levels. + * @throws RuntimeException If this is not an answer to a score question. + */ + public function getScore(): float + { + $this->assertType(ClassificationQuestionTypeEnum::score()); + + return (float) $this->value; + } + + /** + * Gets the model's confidence in the answer. + * + * @since n.e.x.t + * + * @return float|null The confidence between 0 and 1, or null if not reported. + */ + public function getConfidence(): ?float + { + return $this->confidence; + } + + /** + * Gets the probability of each option or level. + * + * @since n.e.x.t + * + * @return array|list The probabilities keyed by option key for a choice question, or + * listed from the lowest level to the highest for a score + * question. Empty if not reported, and always empty for a binary + * question. + */ + public function getProbabilities(): array + { + return $this->probabilities; + } + + /** + * Asserts that this is an answer to a question of the given type. + * + * @since n.e.x.t + * + * @param ClassificationQuestionTypeEnum $type The expected question type. + * @return void + * @throws RuntimeException If this is an answer to a question of another type. + */ + private function assertType(ClassificationQuestionTypeEnum $type): void + { + if (!$this->type->is($type)) { + throw new RuntimeException( + sprintf( + 'This is an answer to a %s question, not a %s question.', + $this->type->value, + $type->value + ) + ); + } + } + + /** + * Validates that a value is a number between 0 and 1. + * + * @since n.e.x.t + * + * @param mixed $value The value to validate. + * @param string $label The label to use in the error message. + * @return float The validated value. + * @throws InvalidArgumentException If the value is not a number between 0 and 1. + */ + private static function validateProbability($value, string $label): float + { + // NaN fails no comparison, so it is excluded explicitly. + if ((!is_int($value) && !is_float($value)) || is_nan((float) $value) || $value < 0 || $value > 1) { + throw new InvalidArgumentException(sprintf('%s must be a number between 0 and 1.', $label)); + } + + return (float) $value; + } + + /** + * {@inheritDoc} + * + * @since n.e.x.t + */ + public static function getJsonSchema(): array + { + $probabilitySchema = [ + 'type' => 'number', + 'minimum' => 0, + 'maximum' => 1, + ]; + + return [ + 'type' => 'object', + 'properties' => [ + self::KEY_TYPE => [ + 'type' => 'string', + 'enum' => ClassificationQuestionTypeEnum::getValues(), + 'description' => 'The type of the question answered.', + ], + self::KEY_VALUE => [ + 'oneOf' => [ + [ + 'type' => 'number', + 'minimum' => 0, + ], + [ + 'type' => 'string', + ], + ], + 'description' => 'The probability for a binary question, the chosen option key for a choice ' + . 'question, or the position on the scale for a score question.', + ], + self::KEY_CONFIDENCE => array_merge($probabilitySchema, [ + 'description' => 'The model\'s confidence in the answer.', + ]), + self::KEY_PROBABILITIES => [ + 'oneOf' => [ + [ + 'type' => 'object', + 'additionalProperties' => $probabilitySchema, + ], + [ + 'type' => 'array', + 'items' => $probabilitySchema, + ], + ], + 'description' => 'The probability of each option keyed by option key for a choice question, ' + . 'or of each level from lowest to highest for a score question.', + ], + ], + 'required' => [self::KEY_TYPE, self::KEY_VALUE], + ]; + } + + /** + * {@inheritDoc} + * + * @since n.e.x.t + * + * @return ClassificationAnswerArrayShape + */ + public function toArray(): array + { + $data = [ + self::KEY_TYPE => $this->type->value, + self::KEY_VALUE => $this->value, + ]; + + if ($this->confidence !== null) { + $data[self::KEY_CONFIDENCE] = $this->confidence; + } + + if ($this->probabilities !== []) { + $data[self::KEY_PROBABILITIES] = $this->probabilities; + } + + return $data; + } + + /** + * {@inheritDoc} + * + * @since n.e.x.t + */ + public static function fromArray(array $array): self + { + static::validateFromArrayData($array, [self::KEY_TYPE, self::KEY_VALUE]); + + return new self( + ClassificationQuestionTypeEnum::from($array[self::KEY_TYPE]), + $array[self::KEY_VALUE], + $array[self::KEY_CONFIDENCE] ?? null, + $array[self::KEY_PROBABILITIES] ?? [] + ); + } +} diff --git a/src/Results/DTO/ClassificationResult.php b/src/Results/DTO/ClassificationResult.php new file mode 100644 index 00000000..2d9a27e6 --- /dev/null +++ b/src/Results/DTO/ClassificationResult.php @@ -0,0 +1,281 @@ +, + * tokenUsage: TokenUsageArrayShape, + * providerMetadata: ProviderMetadataArrayShape, + * modelMetadata: ModelMetadataArrayShape, + * additionalData?: array + * } + * + * @extends AbstractDataTransferObject + */ +class ClassificationResult extends AbstractDataTransferObject implements ResultInterface +{ + public const KEY_ID = 'id'; + public const KEY_ANSWERS = 'answers'; + public const KEY_TOKEN_USAGE = 'tokenUsage'; + public const KEY_PROVIDER_METADATA = 'providerMetadata'; + public const KEY_MODEL_METADATA = 'modelMetadata'; + public const KEY_ADDITIONAL_DATA = 'additionalData'; + + /** + * @var string Unique identifier for this result. + */ + private string $id; + + /** + * @var array The answers, keyed by question key. + */ + private array $answers; + + /** + * @var TokenUsage Token usage statistics. + */ + private TokenUsage $tokenUsage; + + /** + * @var ProviderMetadata Provider metadata. + */ + private ProviderMetadata $providerMetadata; + + /** + * @var ModelMetadata Model metadata. + */ + private ModelMetadata $modelMetadata; + + /** + * @var array Additional data. + */ + private array $additionalData; + + /** + * Constructor. + * + * @since n.e.x.t + * + * @param string $id Unique identifier for this result. + * @param array $answers The answers, keyed by question key. + * @param TokenUsage $tokenUsage Token usage statistics. + * @param ProviderMetadata $providerMetadata Provider metadata. + * @param ModelMetadata $modelMetadata Model metadata. + * @param array $additionalData Additional data. + * + * @throws InvalidArgumentException If no answers are provided. + */ + public function __construct( + string $id, + array $answers, + TokenUsage $tokenUsage, + ProviderMetadata $providerMetadata, + ModelMetadata $modelMetadata, + array $additionalData = [] + ) { + if ($answers === []) { + throw new InvalidArgumentException('At least one answer must be provided.'); + } + + $this->id = $id; + $this->answers = $answers; + $this->tokenUsage = $tokenUsage; + $this->providerMetadata = $providerMetadata; + $this->modelMetadata = $modelMetadata; + $this->additionalData = $additionalData; + } + + /** + * {@inheritDoc} + * + * @since n.e.x.t + */ + public function getId(): string + { + return $this->id; + } + + /** + * Gets the answers. + * + * @since n.e.x.t + * + * @return array The answers, keyed by question key. + */ + public function getAnswers(): array + { + return $this->answers; + } + + /** + * Gets the answer to a question. + * + * @since n.e.x.t + * + * @param string $questionKey The key of the question. + * @return ClassificationAnswer The answer to the question. + * @throws InvalidArgumentException If there is no answer for the question. + */ + public function getAnswer(string $questionKey): ClassificationAnswer + { + if (!isset($this->answers[$questionKey])) { + throw new InvalidArgumentException( + sprintf('No answer found for question "%s".', $questionKey) + ); + } + + return $this->answers[$questionKey]; + } + + /** + * {@inheritDoc} + * + * @since n.e.x.t + */ + public function getTokenUsage(): TokenUsage + { + return $this->tokenUsage; + } + + /** + * {@inheritDoc} + * + * @since n.e.x.t + */ + public function getProviderMetadata(): ProviderMetadata + { + return $this->providerMetadata; + } + + /** + * {@inheritDoc} + * + * @since n.e.x.t + */ + public function getModelMetadata(): ModelMetadata + { + return $this->modelMetadata; + } + + /** + * {@inheritDoc} + * + * @since n.e.x.t + */ + public function getAdditionalData(): array + { + return $this->additionalData; + } + + /** + * {@inheritDoc} + * + * @since n.e.x.t + */ + public static function getJsonSchema(): array + { + return [ + 'type' => 'object', + 'properties' => [ + self::KEY_ID => [ + 'type' => 'string', + 'description' => 'Unique identifier for this result.', + ], + self::KEY_ANSWERS => [ + 'type' => 'object', + 'additionalProperties' => ClassificationAnswer::getJsonSchema(), + 'description' => 'The answers, keyed by question key.', + ], + self::KEY_TOKEN_USAGE => TokenUsage::getJsonSchema(), + self::KEY_PROVIDER_METADATA => ProviderMetadata::getJsonSchema(), + self::KEY_MODEL_METADATA => ModelMetadata::getJsonSchema(), + self::KEY_ADDITIONAL_DATA => [ + 'type' => 'object', + 'additionalProperties' => true, + 'description' => 'Additional provider-specific data.', + ], + ], + 'required' => [ + self::KEY_ID, + self::KEY_ANSWERS, + self::KEY_TOKEN_USAGE, + self::KEY_PROVIDER_METADATA, + self::KEY_MODEL_METADATA, + ], + ]; + } + + /** + * {@inheritDoc} + * + * @since n.e.x.t + * + * @return ClassificationResultArrayShape + */ + public function toArray(): array + { + $data = [ + self::KEY_ID => $this->id, + self::KEY_ANSWERS => array_map( + static fn (ClassificationAnswer $answer): array => $answer->toArray(), + $this->answers + ), + self::KEY_TOKEN_USAGE => $this->tokenUsage->toArray(), + self::KEY_PROVIDER_METADATA => $this->providerMetadata->toArray(), + self::KEY_MODEL_METADATA => $this->modelMetadata->toArray(), + ]; + + if (!empty($this->additionalData)) { + $data[self::KEY_ADDITIONAL_DATA] = $this->additionalData; + } + + return $data; + } + + /** + * {@inheritDoc} + * + * @since n.e.x.t + */ + public static function fromArray(array $array): self + { + static::validateFromArrayData($array, [ + self::KEY_ID, + self::KEY_ANSWERS, + self::KEY_TOKEN_USAGE, + self::KEY_PROVIDER_METADATA, + self::KEY_MODEL_METADATA, + ]); + + return new self( + $array[self::KEY_ID], + array_map( + static fn (array $answer): ClassificationAnswer => ClassificationAnswer::fromArray($answer), + $array[self::KEY_ANSWERS] + ), + TokenUsage::fromArray($array[self::KEY_TOKEN_USAGE]), + ProviderMetadata::fromArray($array[self::KEY_PROVIDER_METADATA]), + ModelMetadata::fromArray($array[self::KEY_MODEL_METADATA]), + $array[self::KEY_ADDITIONAL_DATA] ?? [] + ); + } +} diff --git a/tests/traits/MockModelCreationTrait.php b/tests/traits/MockModelCreationTrait.php index e4fb2c80..0e825b01 100644 --- a/tests/traits/MockModelCreationTrait.php +++ b/tests/traits/MockModelCreationTrait.php @@ -9,6 +9,9 @@ use WordPress\AiClient\Messages\Enums\ModalityEnum; use WordPress\AiClient\Providers\DTO\ProviderMetadata; use WordPress\AiClient\Providers\Enums\ProviderTypeEnum; +use WordPress\AiClient\Providers\Models\Classification\Contracts\ClassificationModelInterface; +use WordPress\AiClient\Providers\Models\Classification\DTO\ClassificationQuestion; +use WordPress\AiClient\Providers\Models\Classification\Enums\ClassificationQuestionTypeEnum; use WordPress\AiClient\Providers\Models\Contracts\ModelInterface; use WordPress\AiClient\Providers\Models\DTO\ModelConfig; use WordPress\AiClient\Providers\Models\DTO\ModelMetadata; @@ -21,6 +24,8 @@ use WordPress\AiClient\Providers\Models\VideoGeneration\Contracts\VideoGenerationModelInterface; use WordPress\AiClient\Providers\ProviderRegistry; use WordPress\AiClient\Results\DTO\Candidate; +use WordPress\AiClient\Results\DTO\ClassificationAnswer; +use WordPress\AiClient\Results\DTO\ClassificationResult; use WordPress\AiClient\Results\DTO\EmbeddingResult; use WordPress\AiClient\Results\DTO\GenerativeAiResult; use WordPress\AiClient\Results\DTO\TokenUsage; @@ -113,6 +118,31 @@ protected function createTestEmbeddingResult(?array $embeddings = null): Embeddi ); } + /** + * Creates a test ClassificationResult for testing purposes. + * + * @param array|null $answers Optional answers for the response. + * @return ClassificationResult + */ + protected function createTestClassificationResult(?array $answers = null): ClassificationResult + { + $answers = $answers ?? ['spam' => new ClassificationAnswer(ClassificationQuestionTypeEnum::binary(), 0.03)]; + + $providerMetadata = new ProviderMetadata( + 'mock', + 'Mock Provider', + ProviderTypeEnum::cloud() + ); + + return new ClassificationResult( + 'test-classification-result-id', + $answers, + new TokenUsage(10, 1, 11), + $providerMetadata, + $this->createTestClassificationModelMetadata() + ); + } + /** * Creates a test model metadata instance for text generation. * @@ -170,6 +200,25 @@ protected function createTestVideoModelMetadata( ); } + /** + * Creates a test model metadata instance for classification. + * + * @param string $id Optional model ID. + * @param string $name Optional model name. + * @return ModelMetadata + */ + protected function createTestClassificationModelMetadata( + string $id = 'test-classification-model', + string $name = 'Test Classification Model' + ): ModelMetadata { + return new ModelMetadata( + $id, + $name, + [CapabilityEnum::classification()], + [] + ); + } + /** * Creates a test model metadata instance for embedding generation. * @@ -472,6 +521,85 @@ public function generateEmbeddingResult(array $inputs): EmbeddingResult }; } + /** + * Creates a mock classification model using anonymous class. + * + * @param ClassificationResult $result The result to return from classification. + * @param ModelMetadata|null $metadata Optional metadata (uses default if not provided). + * @return ModelInterface&ClassificationModelInterface The mock model. + */ + protected function createMockClassificationModel( + ClassificationResult $result, + ?ModelMetadata $metadata = null + ): ModelInterface { + $metadata = $metadata ?? $this->createTestClassificationModelMetadata(); + + $providerMetadata = new ProviderMetadata( + 'mock', + 'Mock Provider', + ProviderTypeEnum::cloud() + ); + + return new class ( + $metadata, + $providerMetadata, + $result + ) implements ModelInterface, ClassificationModelInterface { + private ModelMetadata $metadata; + private ProviderMetadata $providerMetadata; + private ClassificationResult $result; + private ModelConfig $config; + + /** + * @var array{0: array, 1: array}|null The state and + * questions received by the last call to classifyResult(). + */ + public ?array $lastCall = null; + + public function __construct( + ModelMetadata $metadata, + ProviderMetadata $providerMetadata, + ClassificationResult $result + ) { + $this->metadata = $metadata; + $this->providerMetadata = $providerMetadata; + $this->result = $result; + $this->config = new ModelConfig(); + } + + public function metadata(): ModelMetadata + { + return $this->metadata; + } + + public function providerMetadata(): ProviderMetadata + { + return $this->providerMetadata; + } + + public function setConfig(ModelConfig $config): void + { + $this->config = $config; + } + + public function getConfig(): ModelConfig + { + return $this->config; + } + + /** + * @param array $state The state to classify. + * @param array $questions The questions to answer. + */ + public function classifyResult(array $state, array $questions): ClassificationResult + { + $this->lastCall = [$state, $questions]; + + return $this->result; + } + }; + } + /** * Creates a mock model that doesn't implement any generation interfaces. * diff --git a/tests/unit/AiClientTest.php b/tests/unit/AiClientTest.php index c19c882d..ee4f77af 100644 --- a/tests/unit/AiClientTest.php +++ b/tests/unit/AiClientTest.php @@ -7,13 +7,19 @@ use PHPUnit\Framework\TestCase; use RuntimeException; use WordPress\AiClient\AiClient; +use WordPress\AiClient\Builders\ClassificationBuilder; use WordPress\AiClient\Builders\EmbeddingBuilder; use WordPress\AiClient\Common\Exception\InvalidArgumentException; +use WordPress\AiClient\Events\AfterClassifyEvent; +use WordPress\AiClient\Events\BeforeClassifyEvent; use WordPress\AiClient\Messages\DTO\MessagePart; use WordPress\AiClient\Messages\DTO\UserMessage; use WordPress\AiClient\Providers\Contracts\ProviderAvailabilityInterface; +use WordPress\AiClient\Providers\Models\Classification\DTO\ClassificationQuestion; +use WordPress\AiClient\Providers\Models\Classification\Enums\ClassificationQuestionTypeEnum; use WordPress\AiClient\Providers\Models\DTO\ModelConfig; use WordPress\AiClient\Providers\ProviderRegistry; +use WordPress\AiClient\Tests\mocks\MockEventDispatcher; use WordPress\AiClient\Tests\mocks\MockProvider; use WordPress\AiClient\Tests\traits\MockModelCreationTrait; @@ -306,6 +312,57 @@ public function testInputReturnsEmbeddingBuilder(): void $this->assertCount(2, $result); } + /** + * Tests classify() returns a ClassificationBuilder configured with the given state. + */ + public function testClassifyReturnsClassificationBuilder(): void + { + $expectedResult = $this->createTestClassificationResult(); + $mockModel = $this->createMockClassificationModel($expectedResult); + $registry = $this->createRegistryWithMockProvider(); + + $builder = AiClient::classify(['comment' => 'Buy cheap pills!'], $registry); + + $this->assertInstanceOf(ClassificationBuilder::class, $builder); + + $result = $builder + ->withQuestion( + 'spam', + new ClassificationQuestion(ClassificationQuestionTypeEnum::binary(), 'Is this comment spam?') + ) + ->usingModel($mockModel) + ->classifyResult(); + + $this->assertSame($expectedResult, $result); + $this->assertSame(['comment' => 'Buy cheap pills!'], $mockModel->lastCall[0]); + } + + /** + * Tests classify() passes the configured event dispatcher to the builder. + */ + public function testClassifyUsesEventDispatcher(): void + { + $dispatcher = new MockEventDispatcher(); + $mockModel = $this->createMockClassificationModel($this->createTestClassificationResult()); + + AiClient::setEventDispatcher($dispatcher); + + try { + AiClient::classify(['comment' => 'Buy cheap pills!'], $this->createRegistryWithMockProvider()) + ->withQuestion( + 'spam', + new ClassificationQuestion(ClassificationQuestionTypeEnum::binary(), 'Is this comment spam?') + ) + ->usingModel($mockModel) + ->classifyResult(); + } finally { + AiClient::setEventDispatcher(null); + } + + $this->assertCount(1, $dispatcher->getDispatchedEventsOfType(BeforeClassifyEvent::class)); + $this->assertCount(1, $dispatcher->getDispatchedEventsOfType(AfterClassifyEvent::class)); + } + /** * Tests generateTextResult with Message object. diff --git a/tests/unit/Builders/ClassificationBuilderEventDispatchingTest.php b/tests/unit/Builders/ClassificationBuilderEventDispatchingTest.php new file mode 100644 index 00000000..05f852a9 --- /dev/null +++ b/tests/unit/Builders/ClassificationBuilderEventDispatchingTest.php @@ -0,0 +1,111 @@ +registry = new ProviderRegistry(); + $this->registry->registerProvider(MockProvider::class); + $this->dispatcher = new MockEventDispatcher(); + } + + /** + * Tests that events are dispatched for classification. + * + * @return void + */ + public function testEventsAreDispatchedForClassification(): void + { + $result = $this->createTestClassificationResult(); + $model = $this->createMockClassificationModel($result); + $question = new ClassificationQuestion(ClassificationQuestionTypeEnum::binary(), 'Is this comment spam?'); + + $builder = new ClassificationBuilder($this->registry, ['comment' => 'Buy cheap pills!'], $this->dispatcher); + $builder->withQuestion('spam', $question)->usingModel($model); + + $returnedResult = $builder->classifyResult(); + + $beforeEvents = $this->dispatcher->getDispatchedEventsOfType(BeforeClassifyEvent::class); + $afterEvents = $this->dispatcher->getDispatchedEventsOfType(AfterClassifyEvent::class); + + $this->assertCount(1, $beforeEvents); + $this->assertCount(1, $afterEvents); + $this->assertSame(['comment' => 'Buy cheap pills!'], $beforeEvents[0]->getState()); + $this->assertSame(['spam' => $question], $beforeEvents[0]->getQuestions()); + $this->assertSame($model, $beforeEvents[0]->getModel()); + $this->assertEquals(CapabilityEnum::classification(), $beforeEvents[0]->getCapability()); + $this->assertSame(['comment' => 'Buy cheap pills!'], $afterEvents[0]->getState()); + $this->assertSame(['spam' => $question], $afterEvents[0]->getQuestions()); + $this->assertEquals(CapabilityEnum::classification(), $afterEvents[0]->getCapability()); + $this->assertSame($result, $afterEvents[0]->getResult()); + $this->assertSame($returnedResult, $afterEvents[0]->getResult()); + } + + /** + * Tests that no after event is dispatched when the model gives an invalid answer. + * + * @return void + */ + public function testAfterEventIsNotDispatchedForInvalidAnswers(): void + { + $model = $this->createMockClassificationModel( + $this->createTestClassificationResult([ + 'spam' => new ClassificationAnswer(ClassificationQuestionTypeEnum::choice(), 'yes'), + ]) + ); + + $builder = new ClassificationBuilder($this->registry, ['comment' => 'Buy cheap pills!'], $this->dispatcher); + $builder->withQuestion( + 'spam', + new ClassificationQuestion(ClassificationQuestionTypeEnum::binary(), 'Is this comment spam?') + )->usingModel($model); + + try { + $builder->classifyResult(); + $this->fail('Expected the invalid answer to be rejected.'); + } catch (RuntimeException $e) { + $this->assertCount(1, $this->dispatcher->getDispatchedEventsOfType(BeforeClassifyEvent::class)); + $this->assertCount(0, $this->dispatcher->getDispatchedEventsOfType(AfterClassifyEvent::class)); + } + } +} diff --git a/tests/unit/Builders/ClassificationBuilderTest.php b/tests/unit/Builders/ClassificationBuilderTest.php new file mode 100644 index 00000000..215c8f1c --- /dev/null +++ b/tests/unit/Builders/ClassificationBuilderTest.php @@ -0,0 +1,503 @@ +registry = $this->createMock(ProviderRegistry::class); + } + + /** + * Reads a property from a builder. + * + * @param ClassificationBuilder $builder The builder to inspect. + * @param string $propertyName The property name. + * @return mixed The property value. + */ + private function getBuilderProperty(ClassificationBuilder $builder, string $propertyName) + { + $reflection = new \ReflectionClass($builder); + $property = $reflection->getProperty($propertyName); + $property->setAccessible(true); + + return $property->getValue($builder); + } + + /** + * Creates a binary question. + * + * @return ClassificationQuestion + */ + private function createSpamQuestion(): ClassificationQuestion + { + return new ClassificationQuestion(ClassificationQuestionTypeEnum::binary(), 'Is this comment spam?'); + } + + /** + * Creates a builder with one question of each type, answered by a model with the given answers. + * + * @param array $answers The answers the model returns. + * @return ClassificationBuilder + */ + private function createBuilderWithAnswers(array $answers): ClassificationBuilder + { + $builder = new ClassificationBuilder($this->registry, ['comment' => 'Nice post!']); + + return $builder + ->withQuestion('spam', $this->createSpamQuestion()) + ->withQuestion( + 'route', + new ClassificationQuestion( + ClassificationQuestionTypeEnum::choice(), + 'How should a moderator handle it?', + ['approve' => 'Publish it.', 'hold' => 'Look closer.', 'trash' => 'Remove it.'] + ) + ) + ->withQuestion( + 'tone', + new ClassificationQuestion( + ClassificationQuestionTypeEnum::score(), + 'How civil is it?', + ['Hostile.', 'Neutral.', 'Friendly.'] + ) + ) + ->usingModel($this->createMockClassificationModel($this->createTestClassificationResult($answers))); + } + + /** + * Creates a valid answer to each question asked by createBuilderWithAnswers(). + * + * @return array + */ + private static function createValidAnswers(): array + { + return [ + 'spam' => new ClassificationAnswer(ClassificationQuestionTypeEnum::binary(), 0.03), + 'route' => new ClassificationAnswer( + ClassificationQuestionTypeEnum::choice(), + 'approve', + 0.91, + ['approve' => 0.91, 'hold' => 0.07, 'trash' => 0.02] + ), + // The highest level is a valid position on the scale. + 'tone' => new ClassificationAnswer(ClassificationQuestionTypeEnum::score(), 2, 0.6, [0.1, 0.2, 0.7]), + ]; + } + + /** + * Tests that the constructor accepts initial state. + * + * @return void + */ + public function testConstructorWithState(): void + { + $builder = new ClassificationBuilder($this->registry, ['comment' => 'Nice post!']); + + $this->assertSame(['comment' => 'Nice post!'], $this->getBuilderProperty($builder, 'state')); + } + + /** + * Tests that withState() adds to and overwrites existing state. + * + * @return void + */ + public function testWithStateMergesState(): void + { + $builder = new ClassificationBuilder($this->registry, ['comment' => 'First', 'author' => 'Jane']); + + $builder->withState(['comment' => 'Second', 'post' => 'Hello world']); + + $this->assertSame( + ['comment' => 'Second', 'author' => 'Jane', 'post' => 'Hello world'], + $this->getBuilderProperty($builder, 'state') + ); + } + + /** + * Tests that withState() rejects empty state. + * + * @return void + */ + public function testWithStateRejectsEmptyState(): void + { + $this->expectException(InvalidArgumentException::class); + $this->expectExceptionMessage('Classification state cannot be empty.'); + + new ClassificationBuilder($this->registry, []); + } + + /** + * Tests that withState() rejects state that is not keyed by name. + * + * @dataProvider provideStateWithInvalidKeys + * + * @param array $state The state with invalid keys. + * @return void + */ + public function testWithStateRejectsInvalidKeys(array $state): void + { + $this->expectException(InvalidArgumentException::class); + $this->expectExceptionMessage('Classification state keys must be non-empty, non-integer strings.'); + + (new ClassificationBuilder($this->registry))->withState($state); + } + + /** + * Provides state with keys that are not names. + * + * @return array}> + */ + public function provideStateWithInvalidKeys(): array + { + return [ + 'list' => [['Nice post!']], + 'integer-like key' => [['1' => 'Nice post!']], + 'empty key' => [['' => 'Nice post!']], + ]; + } + + /** + * Tests that withQuestion() replaces a question with the same key. + * + * @return void + */ + public function testWithQuestionReplacesExistingKey(): void + { + $replacement = new ClassificationQuestion(ClassificationQuestionTypeEnum::binary(), 'Is this comment rude?'); + + $builder = new ClassificationBuilder($this->registry); + $builder->withQuestion('spam', $this->createSpamQuestion()); + $builder->withQuestion('spam', $replacement); + + $this->assertSame(['spam' => $replacement], $this->getBuilderProperty($builder, 'questions')); + } + + /** + * Tests that withQuestion() rejects keys that PHP would not keep as strings. + * + * @dataProvider provideInvalidQuestionKeys + * + * @param string $key The invalid key. + * @return void + */ + public function testWithQuestionRejectsInvalidKeys(string $key): void + { + $this->expectException(InvalidArgumentException::class); + $this->expectExceptionMessage('Classification question keys must be non-empty, non-integer strings.'); + + (new ClassificationBuilder($this->registry))->withQuestion($key, $this->createSpamQuestion()); + } + + /** + * Provides question keys that are empty or integer-like. + * + * @return array + */ + public function provideInvalidQuestionKeys(): array + { + return [ + 'empty' => [''], + 'whitespace' => [' '], + 'zero' => ['0'], + 'positive integer' => ['12'], + 'negative integer' => ['-3'], + ]; + } + + /** + * Tests that withQuestion() accepts numeric-looking keys that PHP keeps as strings. + * + * @return void + */ + public function testWithQuestionAcceptsNonIntegerNumericKeys(): void + { + $builder = new ClassificationBuilder($this->registry); + $builder->withQuestion('01', $this->createSpamQuestion()); + $builder->withQuestion('1.5', $this->createSpamQuestion()); + + $this->assertSame(['01', '1.5'], array_keys($this->getBuilderProperty($builder, 'questions'))); + } + + /** + * Tests classifyResult() with an explicitly set model. + * + * @return void + */ + public function testClassifyResultWithModel(): void + { + $result = $this->createTestClassificationResult(); + $model = $this->createMockClassificationModel($result); + $question = $this->createSpamQuestion(); + + $this->registry->expects($this->once()) + ->method('bindModelDependencies') + ->with($model); + + $builder = new ClassificationBuilder($this->registry, ['comment' => 'Buy cheap pills!']); + $builder->withQuestion('spam', $question)->usingModel($model); + + $this->assertSame($result, $builder->classifyResult()); + $this->assertSame([['comment' => 'Buy cheap pills!'], ['spam' => $question]], $model->lastCall); + } + + /** + * Tests that classifyResult() discovers a classification model when none is set. + * + * @return void + */ + public function testClassifyResultDiscoversModel(): void + { + $result = $this->createTestClassificationResult(); + $model = $this->createMockClassificationModel($result); + + $this->registry->expects($this->once()) + ->method('findModelsMetadataForSupport') + ->with($this->callback(static function (ModelRequirements $requirements): bool { + return $requirements->getRequiredCapabilities() == [CapabilityEnum::classification()]; + })) + ->willReturn([new ProviderModelsMetadata($model->providerMetadata(), [$model->metadata()])]); + + $this->registry->expects($this->once()) + ->method('getProviderModel') + ->with('mock', 'test-classification-model', $this->isInstanceOf(ModelConfig::class)) + ->willReturn($model); + + $builder = new ClassificationBuilder($this->registry, ['comment' => 'Buy cheap pills!']); + $builder->withQuestion('spam', $this->createSpamQuestion()); + + $this->assertSame($result, $builder->classifyResult()); + } + + /** + * Tests that classifyResult() accepts a valid answer to a question of each type. + * + * @return void + */ + public function testClassifyResultAcceptsValidAnswers(): void + { + $answers = self::createValidAnswers(); + + $result = $this->createBuilderWithAnswers($answers)->classifyResult(); + + $this->assertSame($answers, $result->getAnswers()); + $this->assertSame(0.03, $result->getAnswer('spam')->getProbability()); + $this->assertSame('approve', $result->getAnswer('route')->getChoice()); + $this->assertSame(2.0, $result->getAnswer('tone')->getScore()); + } + + /** + * Tests that classifyResult() requires state. + * + * @return void + */ + public function testClassifyResultRequiresState(): void + { + $builder = new ClassificationBuilder($this->registry); + $builder->withQuestion('spam', $this->createSpamQuestion()); + + $this->expectException(InvalidArgumentException::class); + $this->expectExceptionMessage('Cannot classify empty state. Add state using withState().'); + + $builder->classifyResult(); + } + + /** + * Tests that classifyResult() requires questions. + * + * @return void + */ + public function testClassifyResultRequiresQuestions(): void + { + $builder = new ClassificationBuilder($this->registry, ['comment' => 'Nice post!']); + + $this->expectException(InvalidArgumentException::class); + $this->expectExceptionMessage('Cannot classify without questions. Add questions using withQuestion().'); + + $builder->classifyResult(); + } + + /** + * Tests that classifyResult() rejects a model that does not support classification. + * + * @return void + */ + public function testClassifyResultRejectsUnsupportedModel(): void + { + $builder = new ClassificationBuilder($this->registry, ['comment' => 'Nice post!']); + $builder->withQuestion('spam', $this->createSpamQuestion()) + ->usingModel($this->createMockUnsupportedModel('text-only-model')); + + $this->expectException(RuntimeException::class); + $this->expectExceptionMessage('Model "text-only-model" does not support classification.'); + + $builder->classifyResult(); + } + + /** + * Tests that classifyResult() throws when the model leaves a question unanswered. + * + * @return void + */ + public function testClassifyResultRequiresAnswerForEveryQuestion(): void + { + $answers = self::createValidAnswers(); + unset($answers['route'], $answers['tone']); + + $builder = $this->createBuilderWithAnswers($answers); + + $this->expectException(RuntimeException::class); + $this->expectExceptionMessage('The model did not answer the following questions: route, tone.'); + + $builder->classifyResult(); + } + + /** + * Tests that classifyResult() throws when an answer does not fit its question. + * + * @dataProvider provideInvalidAnswers + * + * @param string $key The key of the question answered invalidly. + * @param ClassificationAnswer $answer The invalid answer. + * @param string $message The expected exception message. + * @return void + */ + public function testClassifyResultRejectsInvalidAnswers( + string $key, + ClassificationAnswer $answer, + string $message + ): void { + $answers = self::createValidAnswers(); + $answers[$key] = $answer; + + $builder = $this->createBuilderWithAnswers($answers); + + $this->expectException(RuntimeException::class); + $this->expectExceptionMessage( + sprintf('The model gave an invalid answer to the question "%s". %s', $key, $message) + ); + + $builder->classifyResult(); + } + + /** + * Provides answers that do not fit their question. + * + * @return array + */ + public function provideInvalidAnswers(): array + { + $choice = ClassificationQuestionTypeEnum::choice(); + $score = ClassificationQuestionTypeEnum::score(); + + return [ + 'wrong type' => [ + 'spam', + new ClassificationAnswer($choice, 'approve'), + 'Expected an answer to a binary question, but received an answer to a choice question.', + ], + // A probability of 1 must not be mistaken for the first level of a scale, or vice versa. + 'score answered as binary' => [ + 'tone', + new ClassificationAnswer(ClassificationQuestionTypeEnum::binary(), 1), + 'Expected an answer to a score question, but received an answer to a binary question.', + ], + 'unknown option' => [ + 'route', + new ClassificationAnswer($choice, 'publish'), + '"publish" is not one of its options.', + ], + 'probabilities for unknown options' => [ + 'route', + new ClassificationAnswer($choice, 'approve', null, ['approve' => 0.8, 'delete' => 0.1, 'spam' => 0.1]), + 'It has probabilities for options the question does not have: delete, spam.', + ], + 'score beyond the scale' => [ + 'tone', + new ClassificationAnswer($score, 2.5), + 'The score 2.5 is outside the scale, which runs from 0 to 2.', + ], + 'probabilities for too few levels' => [ + 'tone', + new ClassificationAnswer($score, 1, null, [0.5, 0.5]), + 'It has probabilities for 2 levels, but the question has 3.', + ], + ]; + } + + /** + * Tests isSupported() when no registered model supports classification. + * + * @return void + */ + public function testIsSupportedReturnsFalseWithoutModels(): void + { + $this->registry->method('findModelsMetadataForSupport')->willReturn([]); + + $builder = new ClassificationBuilder($this->registry, ['comment' => 'Nice post!']); + + $this->assertFalse($builder->isSupported()); + } + + /** + * Tests isSupported() with an explicitly set model. + * + * @return void + */ + public function testIsSupportedWithModel(): void + { + $builder = new ClassificationBuilder($this->registry); + + $builder->usingModel($this->createMockClassificationModel($this->createTestClassificationResult())); + $this->assertTrue($builder->isSupported()); + + $builder->usingModel($this->createMockTextGenerationModel($this->createTestResult())); + $this->assertFalse($builder->isSupported()); + } + + /** + * Tests that a cloned builder does not share model selection with the original. + * + * @return void + */ + public function testCloneDoesNotShareModelSelection(): void + { + $this->registry->method('findModelsMetadataForSupport')->willReturn([]); + + $builder = new ClassificationBuilder($this->registry); + $clone = clone $builder; + + $clone->usingModel($this->createMockClassificationModel($this->createTestClassificationResult())); + + $this->assertTrue($clone->isSupported()); + $this->assertFalse($builder->isSupported()); + } +} diff --git a/tests/unit/Events/AfterClassifyEventTest.php b/tests/unit/Events/AfterClassifyEventTest.php new file mode 100644 index 00000000..45890aae --- /dev/null +++ b/tests/unit/Events/AfterClassifyEventTest.php @@ -0,0 +1,67 @@ + 'Buy cheap pills!']; + $questions = ['spam' => new ClassificationQuestion(ClassificationQuestionTypeEnum::binary(), 'Is it spam?')]; + $result = $this->createTestClassificationResult(); + $model = $this->createMockClassificationModel($result); + $capability = CapabilityEnum::classification(); + + $event = new AfterClassifyEvent($state, $questions, $model, $capability, $result); + + $this->assertSame($state, $event->getState()); + $this->assertSame($questions, $event->getQuestions()); + $this->assertSame($model, $event->getModel()); + $this->assertSame($capability, $event->getCapability()); + $this->assertSame($result, $event->getResult()); + } + + /** + * Tests that cloning the event clones its questions and result. + * + * @return void + */ + public function testCloneClonesQuestionsAndResult(): void + { + $question = new ClassificationQuestion(ClassificationQuestionTypeEnum::binary(), 'Is it spam?'); + $result = $this->createTestClassificationResult(); + + $event = new AfterClassifyEvent( + ['comment' => 'Buy cheap pills!'], + ['spam' => $question], + $this->createMockClassificationModel($result), + CapabilityEnum::classification(), + $result + ); + $clone = clone $event; + + $this->assertSame(['spam'], array_keys($clone->getQuestions())); + $this->assertNotSame($question, $clone->getQuestions()['spam']); + $this->assertNotSame($result, $clone->getResult()); + $this->assertEquals($result, $clone->getResult()); + } +} diff --git a/tests/unit/Events/BeforeClassifyEventTest.php b/tests/unit/Events/BeforeClassifyEventTest.php new file mode 100644 index 00000000..3ea2e28f --- /dev/null +++ b/tests/unit/Events/BeforeClassifyEventTest.php @@ -0,0 +1,64 @@ + 'Buy cheap pills!']; + $questions = ['spam' => new ClassificationQuestion(ClassificationQuestionTypeEnum::binary(), 'Is it spam?')]; + $model = $this->createMockClassificationModel($this->createTestClassificationResult()); + $capability = CapabilityEnum::classification(); + + $event = new BeforeClassifyEvent($state, $questions, $model, $capability); + + $this->assertSame($state, $event->getState()); + $this->assertSame($questions, $event->getQuestions()); + $this->assertSame($model, $event->getModel()); + $this->assertSame($capability, $event->getCapability()); + } + + /** + * Tests that cloning the event clones its questions but keeps their keys. + * + * @return void + */ + public function testCloneClonesQuestions(): void + { + $question = new ClassificationQuestion(ClassificationQuestionTypeEnum::binary(), 'Is it spam?'); + $model = $this->createMockClassificationModel($this->createTestClassificationResult()); + + $event = new BeforeClassifyEvent( + ['comment' => 'Buy cheap pills!'], + ['spam' => $question], + $model, + CapabilityEnum::classification() + ); + $clone = clone $event; + + $this->assertSame(['spam'], array_keys($clone->getQuestions())); + $this->assertNotSame($question, $clone->getQuestions()['spam']); + $this->assertEquals($question, $clone->getQuestions()['spam']); + $this->assertSame($model, $clone->getModel()); + } +} diff --git a/tests/unit/Providers/Models/Classification/DTO/ClassificationQuestionTest.php b/tests/unit/Providers/Models/Classification/DTO/ClassificationQuestionTest.php new file mode 100644 index 00000000..a7050834 --- /dev/null +++ b/tests/unit/Providers/Models/Classification/DTO/ClassificationQuestionTest.php @@ -0,0 +1,338 @@ +assertTrue($question->getType()->isBinary()); + $this->assertSame('Is this comment spam?', $question->getInstructions()); + $this->assertSame([], $question->getCriteria()); + $this->assertSame( + [ + ClassificationQuestion::KEY_TYPE => 'binary', + ClassificationQuestion::KEY_INSTRUCTIONS => 'Is this comment spam?', + ], + $question->toArray() + ); + } + + /** + * Tests creating a choice question. + * + * @return void + */ + public function testCreateChoiceQuestion(): void + { + $criteria = [ + 'approve' => 'The comment is fine to publish.', + 'hold' => 'The comment needs a closer look.', + 'trash' => 'The comment should be removed.', + ]; + + $question = new ClassificationQuestion( + ClassificationQuestionTypeEnum::choice(), + 'How should a moderator handle it?', + $criteria + ); + + $this->assertTrue($question->getType()->isChoice()); + $this->assertSame($criteria, $question->getCriteria()); + $this->assertSame($criteria, $question->toArray()[ClassificationQuestion::KEY_CRITERIA]); + } + + /** + * Tests creating a score question. + * + * @return void + */ + public function testCreateScoreQuestion(): void + { + $levels = ['Not toxic.', 'Somewhat toxic.', 'Very toxic.']; + + $question = new ClassificationQuestion( + ClassificationQuestionTypeEnum::score(), + 'How toxic is the comment?', + $levels + ); + + $this->assertTrue($question->getType()->isScore()); + $this->assertSame($levels, $question->getCriteria()); + $this->assertSame($levels, $question->toArray()[ClassificationQuestion::KEY_CRITERIA]); + } + + /** + * Tests array round trip, including through JSON. + * + * @dataProvider provideQuestions + * + * @param ClassificationQuestion $question The question. + * @return void + */ + public function testArrayRoundTrip(ClassificationQuestion $question): void + { + $this->assertEquals($question, ClassificationQuestion::fromArray($question->toArray())); + + $decoded = json_decode((string) json_encode($question->toArray()), true); + $this->assertEquals($question, ClassificationQuestion::fromArray($decoded)); + } + + /** + * Provides a question of each type. + * + * @return array + */ + public function provideQuestions(): array + { + return [ + 'binary' => [new ClassificationQuestion(ClassificationQuestionTypeEnum::binary(), 'Is it spam?')], + 'choice' => [ + new ClassificationQuestion( + ClassificationQuestionTypeEnum::choice(), + 'How should a moderator handle it?', + ['approve' => 'Publish it.', 'trash' => 'Remove it.'] + ), + ], + 'score' => [ + new ClassificationQuestion( + ClassificationQuestionTypeEnum::score(), + 'How toxic is it?', + ['Not toxic.', 'Very toxic.'] + ), + ], + ]; + } + + /** + * Tests that a choice question serializes its options as a JSON object. + * + * @return void + */ + public function testChoiceCriteriaEncodeAsJsonObject(): void + { + $question = new ClassificationQuestion( + ClassificationQuestionTypeEnum::choice(), + 'How should a moderator handle it?', + ['approve' => 'Publish it.', 'trash' => 'Remove it.'] + ); + + $this->assertStringContainsString( + '"criteria":{"approve":"Publish it.","trash":"Remove it."}', + (string) json_encode($question->toArray()) + ); + } + + /** + * Tests that empty instructions are rejected. + * + * @return void + */ + public function testRejectsEmptyInstructions(): void + { + $this->expectException(InvalidArgumentException::class); + $this->expectExceptionMessage('Classification question instructions cannot be empty.'); + + new ClassificationQuestion(ClassificationQuestionTypeEnum::binary(), ' '); + } + + /** + * Tests that binary questions reject criteria. + * + * @return void + */ + public function testBinaryQuestionRejectsCriteria(): void + { + $this->expectException(InvalidArgumentException::class); + $this->expectExceptionMessage('Binary classification questions do not accept criteria.'); + + new ClassificationQuestion( + ClassificationQuestionTypeEnum::binary(), + 'Is this comment spam?', + ['yes' => 'It is spam.', 'no' => 'It is not spam.'] + ); + } + + /** + * Tests that choice and score questions require at least two criteria. + * + * @dataProvider provideSingleCriterion + * + * @param ClassificationQuestionTypeEnum $type The question type. + * @param array|list $criteria A single criterion of the shape the type takes. + * @return void + */ + public function testRequiresAtLeastTwoCriteria(ClassificationQuestionTypeEnum $type, array $criteria): void + { + $this->expectException(InvalidArgumentException::class); + $this->expectExceptionMessage( + sprintf('Classification questions of type "%s" require at least two criteria.', $type->value) + ); + + new ClassificationQuestion($type, 'Pick one.', $criteria); + } + + /** + * Provides a single criterion for each question type that requires criteria. + * + * @return array|list}> + */ + public function provideSingleCriterion(): array + { + return [ + 'choice' => [ClassificationQuestionTypeEnum::choice(), ['only' => 'The only option.']], + 'score' => [ClassificationQuestionTypeEnum::score(), ['The only level.']], + ]; + } + + /** + * Tests that choice questions reject option keys that are not names. + * + * @dataProvider provideInvalidChoiceCriteria + * + * @param array $criteria The invalid criteria. + * @return void + */ + public function testChoiceQuestionRejectsInvalidOptionKeys(array $criteria): void + { + $this->expectException(InvalidArgumentException::class); + $this->expectExceptionMessage( + 'Choice classification question option keys must be non-empty, non-integer strings.' + ); + + new ClassificationQuestion(ClassificationQuestionTypeEnum::choice(), 'Pick one.', $criteria); + } + + /** + * Provides choice criteria with invalid option keys. + * + * @return array}> + */ + public function provideInvalidChoiceCriteria(): array + { + return [ + 'list' => [['Publish it.', 'Remove it.']], + 'integer-like keys' => [['1' => 'Publish it.', '2' => 'Remove it.']], + 'empty key' => [['' => 'Publish it.', 'trash' => 'Remove it.']], + 'whitespace key' => [[' ' => 'Publish it.', 'trash' => 'Remove it.']], + ]; + } + + /** + * Tests that score questions require a list of levels. + * + * @dataProvider provideInvalidScoreCriteria + * + * @param array $criteria The invalid criteria. + * @return void + */ + public function testScoreQuestionRequiresListOfLevels(array $criteria): void + { + $this->expectException(InvalidArgumentException::class); + $this->expectExceptionMessage( + 'Score classification question criteria must be a list of level descriptions, from lowest to highest.' + ); + + new ClassificationQuestion(ClassificationQuestionTypeEnum::score(), 'How toxic is it?', $criteria); + } + + /** + * Provides score criteria that are not a list. + * + * @return array}> + */ + public function provideInvalidScoreCriteria(): array + { + return [ + 'named levels' => [['low' => 'Not toxic.', 'high' => 'Very toxic.']], + 'levels counted from one' => [[1 => 'Not toxic.', 2 => 'Very toxic.']], + ]; + } + + /** + * Tests that criteria descriptions must be non-empty strings. + * + * @dataProvider provideInvalidDescriptions + * + * @param ClassificationQuestionTypeEnum $type The question type. + * @param array $criteria The criteria with an invalid description. + * @return void + */ + public function testRejectsInvalidCriteriaDescriptions(ClassificationQuestionTypeEnum $type, array $criteria): void + { + $this->expectException(InvalidArgumentException::class); + $this->expectExceptionMessage('Classification question criteria descriptions must be non-empty strings.'); + + new ClassificationQuestion($type, 'Pick one.', $criteria); + } + + /** + * Provides criteria with an invalid description. + * + * @return array}> + */ + public function provideInvalidDescriptions(): array + { + $choice = ClassificationQuestionTypeEnum::choice(); + $score = ClassificationQuestionTypeEnum::score(); + + return [ + 'non-string option' => [$choice, ['approve' => 'Publish it.', 'trash' => 1]], + 'empty option' => [$choice, ['approve' => 'Publish it.', 'trash' => ' ']], + 'non-string level' => [$score, ['Not toxic.', null]], + 'empty level' => [$score, ['Not toxic.', '']], + ]; + } + + /** + * Tests that fromArray() rejects data without required keys. + * + * @return void + */ + public function testFromArrayRequiresTypeAndInstructions(): void + { + $this->assertFalse(ClassificationQuestion::isArrayShape(['type' => 'binary'])); + $this->assertFalse(ClassificationQuestion::isArrayShape(['type' => 'unknown', 'instructions' => 'Why?'])); + $this->assertTrue(ClassificationQuestion::isArrayShape(['type' => 'binary', 'instructions' => 'Why?'])); + } + + /** + * Tests the JSON schema. + * + * @return void + */ + public function testJsonSchema(): void + { + $schema = ClassificationQuestion::getJsonSchema(); + + $this->assertSame( + ['binary', 'choice', 'score'], + $schema['properties'][ClassificationQuestion::KEY_TYPE]['enum'] + ); + $this->assertSame( + ['object', 'array'], + array_column($schema['properties'][ClassificationQuestion::KEY_CRITERIA]['oneOf'], 'type') + ); + $this->assertSame( + [ClassificationQuestion::KEY_TYPE, ClassificationQuestion::KEY_INSTRUCTIONS], + $schema['required'] + ); + } +} diff --git a/tests/unit/Providers/Models/Classification/Enums/ClassificationQuestionTypeEnumTest.php b/tests/unit/Providers/Models/Classification/Enums/ClassificationQuestionTypeEnumTest.php new file mode 100644 index 00000000..ff808323 --- /dev/null +++ b/tests/unit/Providers/Models/Classification/Enums/ClassificationQuestionTypeEnumTest.php @@ -0,0 +1,61 @@ + 'binary', + 'CHOICE' => 'choice', + 'SCORE' => 'score', + ]; + } + + /** + * Tests the specific enum methods. + * + * @return void + */ + public function testSpecificEnumMethods(): void + { + $binary = ClassificationQuestionTypeEnum::binary(); + $this->assertTrue($binary->isBinary()); + $this->assertFalse($binary->isChoice()); + + $choice = ClassificationQuestionTypeEnum::choice(); + $this->assertTrue($choice->isChoice()); + $this->assertFalse($choice->isScore()); + + $score = ClassificationQuestionTypeEnum::score(); + $this->assertTrue($score->isScore()); + $this->assertFalse($score->isBinary()); + } +} diff --git a/tests/unit/Providers/Models/Enums/CapabilityEnumTest.php b/tests/unit/Providers/Models/Enums/CapabilityEnumTest.php index c0f006da..8a83c1cd 100644 --- a/tests/unit/Providers/Models/Enums/CapabilityEnumTest.php +++ b/tests/unit/Providers/Models/Enums/CapabilityEnumTest.php @@ -40,6 +40,7 @@ protected function getExpectedValues(): array 'MUSIC_GENERATION' => 'music_generation', 'VIDEO_GENERATION' => 'video_generation', 'EMBEDDING_GENERATION' => 'embedding_generation', + 'CLASSIFICATION' => 'classification', 'CHAT_HISTORY' => 'chat_history', ]; } @@ -62,5 +63,9 @@ public function testSpecificEnumMethods(): void $chatHistory = CapabilityEnum::chatHistory(); $this->assertTrue($chatHistory->isChatHistory()); $this->assertFalse($chatHistory->isEmbeddingGeneration()); + + $classification = CapabilityEnum::classification(); + $this->assertTrue($classification->isClassification()); + $this->assertFalse($classification->isEmbeddingGeneration()); } } diff --git a/tests/unit/Results/DTO/ClassificationAnswerTest.php b/tests/unit/Results/DTO/ClassificationAnswerTest.php new file mode 100644 index 00000000..b63bd476 --- /dev/null +++ b/tests/unit/Results/DTO/ClassificationAnswerTest.php @@ -0,0 +1,348 @@ +assertTrue($answer->getType()->isBinary()); + $this->assertSame(0.03, $answer->getProbability()); + $this->assertNull($answer->getConfidence()); + $this->assertSame([], $answer->getProbabilities()); + $this->assertSame( + [ + ClassificationAnswer::KEY_TYPE => 'binary', + ClassificationAnswer::KEY_VALUE => 0.03, + ], + $answer->toArray() + ); + } + + /** + * Tests an answer to a choice question. + * + * @return void + */ + public function testChoiceAnswer(): void + { + $probabilities = ['approve' => 0.91, 'hold' => 0.07, 'trash' => 0.02]; + $answer = new ClassificationAnswer(ClassificationQuestionTypeEnum::choice(), 'approve', 0.91, $probabilities); + + $this->assertTrue($answer->getType()->isChoice()); + $this->assertSame('approve', $answer->getChoice()); + $this->assertSame(0.91, $answer->getConfidence()); + $this->assertSame($probabilities, $answer->getProbabilities()); + $this->assertSame( + [ + ClassificationAnswer::KEY_TYPE => 'choice', + ClassificationAnswer::KEY_VALUE => 'approve', + ClassificationAnswer::KEY_CONFIDENCE => 0.91, + ClassificationAnswer::KEY_PROBABILITIES => $probabilities, + ], + $answer->toArray() + ); + } + + /** + * Tests an answer to a score question, including a position between levels. + * + * @return void + */ + public function testScoreAnswer(): void + { + $answer = new ClassificationAnswer(ClassificationQuestionTypeEnum::score(), 2.3, 0.26, [0.05, 0.1, 0.35, 0.5]); + + $this->assertTrue($answer->getType()->isScore()); + $this->assertSame(2.3, $answer->getScore()); + $this->assertSame(0.26, $answer->getConfidence()); + $this->assertSame([0.05, 0.1, 0.35, 0.5], $answer->getProbabilities()); + } + + /** + * Tests that integer values are normalized to floats. + * + * @return void + */ + public function testNormalizesIntegers(): void + { + $binary = new ClassificationAnswer(ClassificationQuestionTypeEnum::binary(), 1, 1); + $score = new ClassificationAnswer(ClassificationQuestionTypeEnum::score(), 3, null, [0, 0, 0, 1]); + + $this->assertSame(1.0, $binary->getProbability()); + $this->assertSame(1.0, $binary->getConfidence()); + $this->assertSame(3.0, $score->getScore()); + $this->assertSame([0.0, 0.0, 0.0, 1.0], $score->getProbabilities()); + } + + /** + * Tests array round trip, including through JSON. + * + * @dataProvider provideAnswers + * + * @param ClassificationAnswer $answer The answer. + * @return void + */ + public function testArrayRoundTrip(ClassificationAnswer $answer): void + { + $this->assertEquals($answer, ClassificationAnswer::fromArray($answer->toArray())); + + $decoded = json_decode((string) json_encode($answer->toArray()), true); + $this->assertEquals($answer, ClassificationAnswer::fromArray($decoded)); + } + + /** + * Provides an answer of each type. + * + * @return array + */ + public function provideAnswers(): array + { + return [ + 'binary' => [new ClassificationAnswer(ClassificationQuestionTypeEnum::binary(), 0.03)], + 'choice' => [ + new ClassificationAnswer( + ClassificationQuestionTypeEnum::choice(), + 'hold', + 0.6, + ['approve' => 0.4, 'hold' => 0.6] + ), + ], + 'score' => [new ClassificationAnswer(ClassificationQuestionTypeEnum::score(), 1.4, 0.5, [0.1, 0.4, 0.5])], + ]; + } + + /** + * Tests that values that do not suit the question type are rejected. + * + * @dataProvider provideInvalidValues + * + * @param ClassificationQuestionTypeEnum $type The question type. + * @param mixed $value The invalid value. + * @param string $message The expected exception message. + * @return void + */ + public function testRejectsInvalidValues(ClassificationQuestionTypeEnum $type, $value, string $message): void + { + $this->expectException(InvalidArgumentException::class); + $this->expectExceptionMessage($message); + + new ClassificationAnswer($type, $value); + } + + /** + * Provides values that do not suit the question type. + * + * @return array + */ + public function provideInvalidValues(): array + { + $binary = ClassificationQuestionTypeEnum::binary(); + $choice = ClassificationQuestionTypeEnum::choice(); + $score = ClassificationQuestionTypeEnum::score(); + + $probabilityMessage = 'Classification answer probability must be a number between 0 and 1.'; + $choiceMessage = 'Classification answer choice must be a non-empty option key.'; + $scoreMessage = 'Classification answer score must be a non-negative number.'; + + return [ + 'binary below zero' => [$binary, -0.1, $probabilityMessage], + 'binary above one' => [$binary, 1.5, $probabilityMessage], + 'binary NaN' => [$binary, NAN, $probabilityMessage], + 'binary option key' => [$binary, 'yes', $probabilityMessage], + 'binary boolean' => [$binary, true, $probabilityMessage], + 'choice number' => [$choice, 1, $choiceMessage], + 'choice empty string' => [$choice, ' ', $choiceMessage], + 'score below zero' => [$score, -1, $scoreMessage], + 'score NaN' => [$score, NAN, $scoreMessage], + 'score infinite' => [$score, INF, $scoreMessage], + 'score level key' => [$score, 'high', $scoreMessage], + ]; + } + + /** + * Tests that an invalid confidence is rejected. + * + * @dataProvider provideInvalidProbabilities + * + * @param float $confidence The invalid confidence. + * @return void + */ + public function testRejectsInvalidConfidence(float $confidence): void + { + $this->expectException(InvalidArgumentException::class); + $this->expectExceptionMessage('Classification answer confidence must be a number between 0 and 1.'); + + new ClassificationAnswer(ClassificationQuestionTypeEnum::choice(), 'approve', $confidence); + } + + /** + * Tests that invalid probabilities are rejected. + * + * @dataProvider provideInvalidProbabilities + * + * @param float $probability The invalid probability. + * @return void + */ + public function testRejectsInvalidProbabilities(float $probability): void + { + $this->expectException(InvalidArgumentException::class); + $this->expectExceptionMessage('Classification answer probabilities must be a number between 0 and 1.'); + + new ClassificationAnswer( + ClassificationQuestionTypeEnum::choice(), + 'approve', + null, + ['approve' => 0.9, 'trash' => $probability] + ); + } + + /** + * Provides numbers that are not valid probabilities. + * + * @return array + */ + public function provideInvalidProbabilities(): array + { + return [ + 'below zero' => [-0.1], + 'above one' => [1.2], + 'NaN' => [NAN], + ]; + } + + /** + * Tests that answers to binary questions reject probabilities. + * + * @return void + */ + public function testBinaryAnswerRejectsProbabilities(): void + { + $this->expectException(InvalidArgumentException::class); + $this->expectExceptionMessage( + 'Classification answers to binary questions do not accept probabilities, as their value is the probability.' + ); + + new ClassificationAnswer(ClassificationQuestionTypeEnum::binary(), 0.9, null, ['true' => 0.9, 'false' => 0.1]); + } + + /** + * Tests that answers to choice questions require probabilities keyed by option key. + * + * @return void + */ + public function testChoiceAnswerRequiresProbabilitiesKeyedByOption(): void + { + $this->expectException(InvalidArgumentException::class); + $this->expectExceptionMessage( + 'Classification answer probabilities for a choice question must be keyed by option key.' + ); + + new ClassificationAnswer(ClassificationQuestionTypeEnum::choice(), 'approve', null, [0.9, 0.1]); + } + + /** + * Tests that answers to score questions require a list of probabilities. + * + * @return void + */ + public function testScoreAnswerRequiresListOfProbabilities(): void + { + $this->expectException(InvalidArgumentException::class); + $this->expectExceptionMessage( + 'Classification answer probabilities for a score question must be a list, ' + . 'from the lowest level to the highest.' + ); + + new ClassificationAnswer(ClassificationQuestionTypeEnum::score(), 1.0, null, ['low' => 0.2, 'high' => 0.8]); + } + + /** + * Tests that reading the value of another question type throws. + * + * @dataProvider provideMismatchedGetters + * + * @param ClassificationAnswer $answer The answer. + * @param string $getter The getter for another question type. + * @param string $message The expected exception message. + * @return void + */ + public function testGetterForAnotherTypeThrows(ClassificationAnswer $answer, string $getter, string $message): void + { + $this->expectException(RuntimeException::class); + $this->expectExceptionMessage($message); + + $answer->$getter(); + } + + /** + * Provides answers with a getter for another question type. + * + * @return array + */ + public function provideMismatchedGetters(): array + { + $binary = new ClassificationAnswer(ClassificationQuestionTypeEnum::binary(), 0.5); + $choice = new ClassificationAnswer(ClassificationQuestionTypeEnum::choice(), 'approve'); + $score = new ClassificationAnswer(ClassificationQuestionTypeEnum::score(), 1.0); + + return [ + 'choice of binary' => [ + $binary, + 'getChoice', + 'This is an answer to a binary question, not a choice question.', + ], + 'score of binary' => [$binary, 'getScore', 'This is an answer to a binary question, not a score question.'], + 'probability of choice' => [ + $choice, + 'getProbability', + 'This is an answer to a choice question, not a binary question.', + ], + 'score of choice' => [$choice, 'getScore', 'This is an answer to a choice question, not a score question.'], + 'probability of score' => [ + $score, + 'getProbability', + 'This is an answer to a score question, not a binary question.', + ], + 'choice of score' => [$score, 'getChoice', 'This is an answer to a score question, not a choice question.'], + ]; + } + + /** + * Tests the JSON schema. + * + * @return void + */ + public function testJsonSchema(): void + { + $schema = ClassificationAnswer::getJsonSchema(); + + $this->assertSame( + ['binary', 'choice', 'score'], + $schema['properties'][ClassificationAnswer::KEY_TYPE]['enum'] + ); + $this->assertSame( + ['object', 'array'], + array_column($schema['properties'][ClassificationAnswer::KEY_PROBABILITIES]['oneOf'], 'type') + ); + $this->assertSame([ClassificationAnswer::KEY_TYPE, ClassificationAnswer::KEY_VALUE], $schema['required']); + } +} diff --git a/tests/unit/Results/DTO/ClassificationResultTest.php b/tests/unit/Results/DTO/ClassificationResultTest.php new file mode 100644 index 00000000..f036d9d3 --- /dev/null +++ b/tests/unit/Results/DTO/ClassificationResultTest.php @@ -0,0 +1,148 @@ + $answers The answers. + * @return ClassificationResult + */ + private function createClassificationResult(array $answers): ClassificationResult + { + return new ClassificationResult( + 'classification-result-id', + $answers, + new TokenUsage(312, 9, 321), + new ProviderMetadata('mock', 'Mock Provider', ProviderTypeEnum::cloud()), + new ModelMetadata( + 'mock-classification-model', + 'Mock Classification Model', + [CapabilityEnum::classification()], + [] + ), + ['providerResultId' => 'provider-123'] + ); + } + + /** + * Tests the getters. + * + * @return void + */ + public function testGetters(): void + { + $spam = new ClassificationAnswer(ClassificationQuestionTypeEnum::binary(), 0.03); + $route = new ClassificationAnswer( + ClassificationQuestionTypeEnum::choice(), + 'approve', + 0.91, + ['approve' => 0.91, 'hold' => 0.07, 'trash' => 0.02] + ); + + $result = $this->createClassificationResult(['spam' => $spam, 'route' => $route]); + + $this->assertSame('classification-result-id', $result->getId()); + $this->assertSame(['spam' => $spam, 'route' => $route], $result->getAnswers()); + $this->assertSame($spam, $result->getAnswer('spam')); + $this->assertSame($route, $result->getAnswer('route')); + $this->assertSame(312, $result->getTokenUsage()->getPromptTokens()); + $this->assertSame('mock', $result->getProviderMetadata()->getId()); + $this->assertSame('mock-classification-model', $result->getModelMetadata()->getId()); + $this->assertSame(['providerResultId' => 'provider-123'], $result->getAdditionalData()); + } + + /** + * Tests array round trip, including through JSON. + * + * @return void + */ + public function testArrayRoundTrip(): void + { + $result = $this->createClassificationResult([ + 'spam' => new ClassificationAnswer(ClassificationQuestionTypeEnum::binary(), 0.03), + 'route' => new ClassificationAnswer( + ClassificationQuestionTypeEnum::choice(), + 'approve', + 0.91, + ['approve' => 0.91, 'trash' => 0.09] + ), + 'tone' => new ClassificationAnswer(ClassificationQuestionTypeEnum::score(), 1.7, 0.8, [0.05, 0.2, 0.75]), + ]); + + $array = $result->toArray(); + + $this->assertSame( + [ClassificationAnswer::KEY_TYPE => 'binary', ClassificationAnswer::KEY_VALUE => 0.03], + $array[ClassificationResult::KEY_ANSWERS]['spam'] + ); + $this->assertEquals($result, ClassificationResult::fromArray($array)); + + $decoded = json_decode((string) json_encode($array), true); + $this->assertEquals($result, ClassificationResult::fromArray($decoded)); + } + + /** + * Tests that at least one answer is required. + * + * @return void + */ + public function testRequiresAtLeastOneAnswer(): void + { + $this->expectException(InvalidArgumentException::class); + $this->expectExceptionMessage('At least one answer must be provided.'); + + $this->createClassificationResult([]); + } + + /** + * Tests that getting an unknown answer throws. + * + * @return void + */ + public function testGetAnswerThrowsForUnknownQuestion(): void + { + $result = $this->createClassificationResult([ + 'spam' => new ClassificationAnswer(ClassificationQuestionTypeEnum::binary(), 0.03), + ]); + + $this->expectException(InvalidArgumentException::class); + $this->expectExceptionMessage('No answer found for question "route".'); + + $result->getAnswer('route'); + } + + /** + * Tests the JSON schema. + * + * @return void + */ + public function testJsonSchema(): void + { + $schema = ClassificationResult::getJsonSchema(); + + $this->assertSame( + ClassificationAnswer::getJsonSchema(), + $schema['properties'][ClassificationResult::KEY_ANSWERS]['additionalProperties'] + ); + $this->assertContains(ClassificationResult::KEY_ANSWERS, $schema['required']); + } +}