Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions src/platform/CHANGELOG.md
Original file line number Diff line number Diff line change
@@ -1,6 +1,11 @@
CHANGELOG
=========

0.14
----

* Add the serving provider to `ResultConvertedEvent` and `ResultErrorEvent`, so listeners can attribute a resolved result to the provider that produced it (e.g. to release held capacity)

0.13
----

Expand Down
7 changes: 7 additions & 0 deletions src/platform/src/Event/ResultConvertedEvent.php
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
namespace Symfony\AI\Platform\Event;

use Symfony\AI\Platform\Model;
use Symfony\AI\Platform\ProviderInterface;
use Symfony\AI\Platform\Result\ResultInterface;
use Symfony\Contracts\EventDispatcher\Event;

Expand All @@ -35,9 +36,15 @@ public function __construct(
private ResultInterface $result,
private readonly array $options = [],
private readonly array|string|object $input = [],
private readonly ?ProviderInterface $provider = null,
) {
}

public function getProvider(): ?ProviderInterface
{
return $this->provider;
}

public function getModel(): Model
{
return $this->model;
Expand Down
7 changes: 7 additions & 0 deletions src/platform/src/Event/ResultErrorEvent.php
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
namespace Symfony\AI\Platform\Event;

use Symfony\AI\Platform\Model;
use Symfony\AI\Platform\ProviderInterface;
use Symfony\Contracts\EventDispatcher\Event;

/**
Expand All @@ -34,9 +35,15 @@ public function __construct(
private readonly \Throwable $error,
private readonly array $options = [],
private readonly array|string|object $input = [],
private readonly ?ProviderInterface $provider = null,
) {
}

public function getProvider(): ?ProviderInterface
{
return $this->provider;
}

public function getModel(): Model
{
return $this->model;
Expand Down
4 changes: 2 additions & 2 deletions src/platform/src/Provider.php
Original file line number Diff line number Diff line change
Expand Up @@ -118,13 +118,13 @@ public function invoke(string|Model $model, array|string|object $input, array $o

if (null !== $this->eventDispatcher) {
$deferredResult->onConvert(function (ResultInterface $result) use ($model, $options, $input): ResultInterface {
$event = new ResultConvertedEvent($model, $result, $options, $input);
$event = new ResultConvertedEvent($model, $result, $options, $input, $this);
$this->eventDispatcher->dispatch($event);

return $event->getResult();
});
$deferredResult->onError(function (\Throwable $error) use ($model, $options, $input): void {
$this->eventDispatcher->dispatch(new ResultErrorEvent($model, $error, $options, $input));
$this->eventDispatcher->dispatch(new ResultErrorEvent($model, $error, $options, $input, $this));
});
}

Expand Down
13 changes: 12 additions & 1 deletion src/platform/tests/Event/ResultConvertedEventTest.php
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
use Symfony\AI\Platform\Capability;
use Symfony\AI\Platform\Event\ResultConvertedEvent;
use Symfony\AI\Platform\Model;
use Symfony\AI\Platform\ProviderInterface;
use Symfony\AI\Platform\Result\TextResult;

final class ResultConvertedEventTest extends TestCase
Expand All @@ -24,13 +25,15 @@ public function testGettersReturnConstructorValues()
$model = new Model('test-model', [Capability::OUTPUT_TEXT]);
$result = new TextResult('Hello');
$options = ['temperature' => 0.7];
$provider = $this->createStub(ProviderInterface::class);

$event = new ResultConvertedEvent($model, $result, $options, 'Hello?');
$event = new ResultConvertedEvent($model, $result, $options, 'Hello?', $provider);

$this->assertSame($model, $event->getModel());
$this->assertSame($result, $event->getResult());
$this->assertSame($options, $event->getOptions());
$this->assertSame('Hello?', $event->getInput());
$this->assertSame($provider, $event->getProvider());
}

public function testSetResultOverridesResolvedResult()
Expand All @@ -42,4 +45,12 @@ public function testSetResultOverridesResolvedResult()

$this->assertSame($newResult, $event->getResult());
}

public function testProviderIsOptional()
{
$model = new Model('test-model', [Capability::OUTPUT_TEXT]);
$result = new TextResult('Hello');

$this->assertNull((new ResultConvertedEvent($model, $result))->getProvider());
}
}
13 changes: 12 additions & 1 deletion src/platform/tests/Event/ResultErrorEventTest.php
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
use Symfony\AI\Platform\Event\ResultErrorEvent;
use Symfony\AI\Platform\Exception\RuntimeException;
use Symfony\AI\Platform\Model;
use Symfony\AI\Platform\ProviderInterface;

final class ResultErrorEventTest extends TestCase
{
Expand All @@ -24,12 +25,22 @@ public function testGettersReturnConstructorValues()
$model = new Model('test-model', [Capability::OUTPUT_TEXT]);
$error = new RuntimeException('conversion failed');
$options = ['temperature' => 0.7];
$provider = $this->createStub(ProviderInterface::class);

$event = new ResultErrorEvent($model, $error, $options, 'Hello?');
$event = new ResultErrorEvent($model, $error, $options, 'Hello?', $provider);

$this->assertSame($model, $event->getModel());
$this->assertSame($error, $event->getError());
$this->assertSame($options, $event->getOptions());
$this->assertSame('Hello?', $event->getInput());
$this->assertSame($provider, $event->getProvider());
}

public function testProviderIsOptional()
{
$model = new Model('test-model', [Capability::OUTPUT_TEXT]);
$error = new RuntimeException('conversion failed');

$this->assertNull((new ResultErrorEvent($model, $error))->getProvider());
}
}
2 changes: 2 additions & 0 deletions src/platform/tests/ProviderTest.php
Original file line number Diff line number Diff line change
Expand Up @@ -253,6 +253,7 @@ public function testInvokeDispatchesResultConvertedEventOnConversion()
$this->assertCount(3, $dispatchedEvents);
$this->assertInstanceOf(ResultConvertedEvent::class, $dispatchedEvents[2]);
$this->assertSame($textResult, $dispatchedEvents[2]->getResult());
$this->assertSame($provider, $dispatchedEvents[2]->getProvider());
}

public function testInvokeDispatchesResultErrorEventOnConversionFailure()
Expand Down Expand Up @@ -293,6 +294,7 @@ public function testInvokeDispatchesResultErrorEventOnConversionFailure()

$this->assertInstanceOf(ResultErrorEvent::class, $dispatchedEvents[2]);
$this->assertSame($exception, $dispatchedEvents[2]->getError());
$this->assertSame($provider, $dispatchedEvents[2]->getProvider());
}

public function testGetModelCatalog()
Expand Down
Loading