Skip to content
Merged
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
3 changes: 2 additions & 1 deletion packages/cli/src/providers.ts
Original file line number Diff line number Diff line change
Expand Up @@ -149,7 +149,8 @@ async function ensurePythonProvidersRegistered(): Promise<void> {
if (listRegisteredProviderIds().includes(info.id)) continue;
registerProvider(info.id, PythonProvider as never, {
_bridge: bridge,
_id: info.id
_id: info.id,
_capabilities: info.capabilities
});
}
})().catch((err) => {
Expand Down
59 changes: 53 additions & 6 deletions packages/runtime/src/providers/python-provider.ts
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@ import type {
ASRModel,
EmbeddingModel,
VideoModel,
MusicModel,
Message,
ProviderTool,
ProviderStreamItem,
Expand All @@ -23,16 +24,21 @@ import type {
ImageToImageParams,
TextToSpeechParams,
TextToVideoParams,
ImageToVideoParams
ImageToVideoParams,
TextToMusicParams,
EncodedAudioResult
} from "./types.js";
import type { PythonBridgeBase } from "../python-bridge-base.js";
import { isRecord, isString } from "@nodetool-ai/protocol";
import { sniffAudioMime } from "./audio-mime.js";

type PythonProviderOptions = Record<string, unknown> & {
_id: string;
_bridge: PythonBridgeBase;
/** Provider id understood by the Python worker when the public id is aliased. */
_bridgeProviderId?: string;
/** Operations advertised by the Python worker for this provider. */
_capabilities?: string[];
};

function parseModelAdapter(value: unknown): ModelAdapterInfo | undefined {
Expand Down Expand Up @@ -71,6 +77,8 @@ export class PythonProvider extends BaseProvider {
private _bridge: PythonBridgeBase;
private _pythonProviderId: string;
private _secrets: Record<string, string>;
private _supportsStreamingTTS = true;
private _supportsEncodedTTS = true;

constructor(
providerId: string,
Expand All @@ -94,11 +102,17 @@ export class PythonProvider extends BaseProvider {
return;
}

const { _id, _bridge, _bridgeProviderId, ...rawSecrets } =
const { _id, _bridge, _bridgeProviderId, _capabilities, ...rawSecrets } =
providerIdOrOptions;
super(_id);
this._bridge = _bridge;
this._pythonProviderId = _bridgeProviderId ?? _id;
if (Array.isArray(_capabilities)) {
this._supportsStreamingTTS = _capabilities.includes("text_to_speech");
this._supportsEncodedTTS = _capabilities.includes(
"text_to_speech_encoded"
);
}
this._secrets = Object.fromEntries(
Object.entries(rawSecrets).filter(
(entry): entry is [string, string] => typeof entry[1] === "string"
Expand Down Expand Up @@ -160,6 +174,10 @@ export class PythonProvider extends BaseProvider {
return this._getModels("video") as Promise<VideoModel[]>;
}

async getAvailableMusicModels(): Promise<MusicModel[]> {
return this._getModels("music") as Promise<MusicModel[]>;
}

private async _getModels(modelType: string): Promise<unknown[]> {
const models = await this._bridge.getProviderModels(
this._pythonProviderId,
Expand Down Expand Up @@ -262,22 +280,26 @@ export class PythonProvider extends BaseProvider {
// ── Media generation ──────────────────────────────────────────────

async textToImage(params: TextToImageParams): Promise<Uint8Array> {
const { signal, ...wireParams } = params;
return this._bridge.providerTextToImage(
this._pythonProviderId,
{ ...params },
this._secrets
{ ...wireParams, model: params.model.id },
this._secrets,
signal
);
}

async imageToImage(
images: Uint8Array[],
params: ImageToImageParams
): Promise<Uint8Array> {
const { signal, ...wireParams } = params;
return this._bridge.providerImageToImage(
this._pythonProviderId,
images[0] ?? new Uint8Array(),
{ ...params },
this._secrets
{ ...wireParams, model: params.model.id },
this._secrets,
signal
);
}

Expand Down Expand Up @@ -333,6 +355,31 @@ export class PythonProvider extends BaseProvider {
}
}

override supportsStreamingTextToSpeech(): boolean {
return this._supportsStreamingTTS;
}

async textToSpeechEncoded(
args: TextToSpeechParams
): Promise<EncodedAudioResult | null> {
if (!this._supportsEncodedTTS) return null;
const data = await this._bridge.providerTTSEncoded(
this._pythonProviderId,
{ ...args },
this._secrets
);
return { data, mimeType: sniffAudioMime(data) };
}

async textToMusic(params: TextToMusicParams): Promise<EncodedAudioResult> {
const data = await this._bridge.providerTextToAudio(
this._pythonProviderId,
{ ...params, model: params.model.id },
this._secrets
);
return { data, mimeType: sniffAudioMime(data) };
}

async automaticSpeechRecognition(args: {
audio: Uint8Array;
model: string;
Expand Down
66 changes: 51 additions & 15 deletions packages/runtime/src/python-bridge-base.ts
Original file line number Diff line number Diff line change
Expand Up @@ -1087,29 +1087,39 @@ export abstract class PythonBridgeBase
async providerTextToImage(
providerId: string,
params: Record<string, unknown>,
secrets?: Record<string, string>
secrets?: Record<string, string>,
signal?: AbortSignal
): Promise<Uint8Array> {
const result = await this._providerCall("provider.text_to_image", {
provider: providerId,
params,
secrets: secrets ?? {}
});
return (result as { blobs: Record<string, Uint8Array> }).blobs.image;
const result = await this._providerBlobCall(
"provider.text_to_image",
{
provider: providerId,
params,
secrets: secrets ?? {}
},
signal
);
return result.blobs.image;
}

async providerImageToImage(
providerId: string,
image: Uint8Array,
params: Record<string, unknown>,
secrets?: Record<string, string>
secrets?: Record<string, string>,
signal?: AbortSignal
): Promise<Uint8Array> {
const result = await this._providerCall("provider.image_to_image", {
provider: providerId,
image,
params,
secrets: secrets ?? {}
});
return (result as { blobs: Record<string, Uint8Array> }).blobs.image;
const result = await this._providerBlobCall(
"provider.image_to_image",
{
provider: providerId,
image,
params,
secrets: secrets ?? {}
},
signal
);
return result.blobs.image;
}

async providerTextToVideo(
Expand Down Expand Up @@ -1150,6 +1160,32 @@ export abstract class PythonBridgeBase
return result.blobs.video;
}

async providerTextToAudio(
providerId: string,
params: Record<string, unknown>,
secrets?: Record<string, string>
): Promise<Uint8Array> {
const result = await this._providerBlobCall("provider.text_to_audio", {
provider: providerId,
params,
secrets: secrets ?? {}
});
return result.blobs.audio;
}

async providerTTSEncoded(
providerId: string,
params: Record<string, unknown>,
secrets?: Record<string, string>
): Promise<Uint8Array> {
const result = await this._providerBlobCall("provider.tts_encoded", {
provider: providerId,
params,
secrets: secrets ?? {}
});
return result.blobs.audio;
}

async *providerTTS(
providerId: string,
text: string,
Expand Down
16 changes: 14 additions & 2 deletions packages/runtime/src/python-bridge-types.ts
Original file line number Diff line number Diff line change
Expand Up @@ -567,13 +567,15 @@ export interface PythonBridge extends EventEmitter {
providerTextToImage(
providerId: string,
params: Record<string, unknown>,
secrets?: Record<string, string>
secrets?: Record<string, string>,
signal?: AbortSignal
): Promise<Uint8Array>;
providerImageToImage(
providerId: string,
image: Uint8Array,
params: Record<string, unknown>,
secrets?: Record<string, string>
secrets?: Record<string, string>,
signal?: AbortSignal
): Promise<Uint8Array>;
providerTextToVideo(
providerId: string,
Expand All @@ -588,6 +590,16 @@ export interface PythonBridge extends EventEmitter {
secrets?: Record<string, string>,
signal?: AbortSignal
): Promise<Uint8Array>;
providerTextToAudio(
providerId: string,
params: Record<string, unknown>,
secrets?: Record<string, string>
): Promise<Uint8Array>;
providerTTSEncoded(
providerId: string,
params: Record<string, unknown>,
secrets?: Record<string, string>
): Promise<Uint8Array>;
providerASR(
providerId: string,
audio: Uint8Array,
Expand Down
27 changes: 23 additions & 4 deletions packages/runtime/src/swappable-python-bridge.ts
Original file line number Diff line number Diff line change
Expand Up @@ -226,22 +226,25 @@ export class SwappableBridge extends EventEmitter implements PythonBridge {
providerTextToImage(
providerId: string,
params: Record<string, unknown>,
secrets?: Record<string, string>
secrets?: Record<string, string>,
signal?: AbortSignal
): Promise<Uint8Array> {
return this._target.providerTextToImage(providerId, params, secrets);
return this._target.providerTextToImage(providerId, params, secrets, signal);
}

providerImageToImage(
providerId: string,
image: Uint8Array,
params: Record<string, unknown>,
secrets?: Record<string, string>
secrets?: Record<string, string>,
signal?: AbortSignal
): Promise<Uint8Array> {
return this._target.providerImageToImage(
providerId,
image,
params,
secrets
secrets,
signal
);
}

Expand Down Expand Up @@ -275,6 +278,22 @@ export class SwappableBridge extends EventEmitter implements PythonBridge {
);
}

providerTextToAudio(
providerId: string,
params: Record<string, unknown>,
secrets?: Record<string, string>
): Promise<Uint8Array> {
return this._target.providerTextToAudio(providerId, params, secrets);
}

providerTTSEncoded(
providerId: string,
params: Record<string, unknown>,
secrets?: Record<string, string>
): Promise<Uint8Array> {
return this._target.providerTTSEncoded(providerId, params, secrets);
}

providerASR(
providerId: string,
audio: Uint8Array,
Expand Down
Loading
Loading