chore: import upstream snapshot with attribution
CI / Run CI (push) Has been cancelled
CI / check-backend (push) Has been cancelled
CI / check-frontend (push) Has been cancelled
CI / tests (push) Has been cancelled
CI / e2e-tests (push) Has been cancelled
Copilot Setup Steps / copilot-setup-steps (push) Has been cancelled

This commit is contained in:
wehub-resource-sync
2026-07-13 12:48:47 +08:00
commit b7f52be4c9
653 changed files with 105877 additions and 0 deletions
+178
View File
@@ -0,0 +1,178 @@
## Overview
The `@chainlit/react-client` package provides a set of React hooks as well as an API client to connect to your [Chainlit](https://github.com/Chainlit/chainlit) application from any React application. The package includes hooks for managing chat sessions, messages, data, and interactions.
## Installation
To install the package, run the following command in your project directory:
```sh
npm install @chainlit/react-client
```
This package use [Recoil](https://github.com/facebookexperimental/Recoil) to manage its state. This means you will have to wrap your application in a recoil provider:
```tsx
import React from 'react';
import ReactDOM from 'react-dom/client';
import { RecoilRoot } from 'recoil';
import { ChainlitAPI, ChainlitContext } from '@chainlit/react-client';
const CHAINLIT_SERVER_URL = 'http://localhost:8000';
const apiClient = new ChainlitAPI(CHAINLIT_SERVER_URL, 'webapp');
ReactDOM.createRoot(document.getElementById('root') as HTMLElement).render(
<React.StrictMode>
<ChainlitContext.Provider value={apiClient}>
<RecoilRoot>
<MyApp />
</RecoilRoot>
</ChainlitContext.Provider>
</React.StrictMode>
);
```
## Usage
### `useChatSession`
This hook is responsible for managing the chat session's connection to the WebSocket server.
#### Methods
- `connect`: Establishes a connection to the WebSocket server.
- `disconnect`: Disconnects from the WebSocket server.
- `setChatProfile`: Sets the chat profile state.
#### Example
```jsx
import { useChatSession } from '@chainlit/react-client';
const ChatComponent = () => {
const { connect, disconnect, chatProfile, setChatProfile } = useChatSession();
// Connect to the WebSocket server
useEffect(() => {
connect({
userEnv: {
/* user environment variables */
}
});
return () => {
disconnect();
};
}, []);
// Rest of your component logic
};
```
### `useChatMessages`
This hook provides access to the chat messages and the first user message.
#### Properties
- `messages`: An array of chat messages.
- `firstUserMessage`: The first message from the user.
#### Example
```jsx
import { useChatMessages } from '@chainlit/react-client';
const MessagesComponent = () => {
const { messages, firstUserMessage } = useChatMessages();
// Render your messages
return (
<div>
{messages.map((message) => (
<p key={message.id}>{message.output}</p>
))}
</div>
);
};
```
### `useChatData`
This hook provides access to various chat-related data and states.
#### Properties
- `actions`: An array of actions.
- `askUser`: The current ask user state.
- `avatars`: An array of avatar elements.
- `chatSettingsDefaultValue`: The default value for chat settings.
- `chatSettingsInputs`: The current chat settings inputs.
- `chatSettingsValue`: The current value of chat settings.
- `connected`: A boolean indicating if the WebSocket connection is established.
- `disabled`: A boolean indicating if the chat is disabled.
- `elements`: An array of chat elements.
- `error`: A boolean indicating if there is an error in the session.
- `loading`: A boolean indicating if the chat is in a loading state.
- `tasklists`: An array of tasklist elements.
#### Example
```jsx
import { useChatData } from '@chainlit/react-client';
const ChatDataComponent = () => {
const { loading, connected, error } = useChatData();
// Use the data to render your component
if (loading) return <p>Loading...</p>;
if (error) return <p>Error connecting to chat...</p>;
if (!connected) return <p>Disconnected...</p>;
// Rest of your component logic
};
```
### `useChatInteract`
This hook provides methods to interact with the chat, such as sending messages, replying, and updating settings.
#### Methods
- `callAction`: Calls an action.
- `clear`: Clears the chat session.
- `replyMessage`: Replies to a message.
- `sendMessage`: Sends a message.
- `stopTask`: Stops the current task.
- `setIdToResume`: Sets the ID to resume a thread.
- `updateChatSettings`: Updates the chat settings.
#### Example
```jsx
import { useChatInteract } from '@chainlit/react-client';
const InteractionComponent = () => {
const { sendMessage, replyMessage } = useChatInteract();
const handleSendMessage = () => {
const message = { output: 'Hello, World!', id: 'message-id' };
sendMessage(message);
};
const handleReplyMessage = () => {
const message = { output: 'Replying to your message', id: 'reply-id' };
replyMessage(message);
};
// Render your interaction component
return (
<div>
<button onClick={handleSendMessage}>Send Message</button>
<button onClick={handleReplyMessage}>Reply to Message</button>
</div>
);
};
```
+76
View File
@@ -0,0 +1,76 @@
{
"name": "@chainlit/react-client",
"description": "Websocket client to connect to your chainlit app.",
"version": "0.4.2",
"scripts": {
"dev": "tsup src/index.ts --clean --format esm,cjs --dts --external react --external recoil --minify --sourcemap --treeshake",
"build": "tsup src/index.ts --tsconfig tsconfig.build.json --clean --format esm,cjs --dts --external react --external recoil --minify --sourcemap --treeshake",
"type-check": "tsc --noemit",
"test": "echo no tests yet"
},
"repository": {
"type": "git",
"url": "git+https://github.com/Chainlit/chainlit.git",
"directory": "libs/react-client"
},
"private": false,
"keywords": [
"llm",
"ai",
"chain of thought"
],
"author": "Chainlit",
"license": "Apache-2.0",
"files": [
"dist",
"README.md"
],
"main": "dist/index.js",
"module": "dist/index.mjs",
"types": "dist/index.d.ts",
"devDependencies": {
"@swc/core": "^1.3.86",
"@testing-library/jest-dom": "^5.17.0",
"@testing-library/react": "^14.0.0",
"@types/lodash": "^4.14.199",
"@types/uuid": "^9.0.3",
"@vitejs/plugin-react": "^4.0.4",
"@vitejs/plugin-react-swc": "^3.3.2",
"jsdom": "^22.1.0",
"tslib": "^2.6.2",
"tsup": "^7.2.0",
"typescript": "^5.2.2",
"vite": "^5.4.14",
"vite-tsconfig-paths": "^4.2.0",
"vitest": "^0.34.4"
},
"peerDependencies": {
"@types/react": "^18.3.1",
"react": "^18.3.1",
"react-dom": "^18.3.1",
"recoil": "^0.7.7"
},
"dependencies": {
"jwt-decode": "^3.1.2",
"lodash": "^4.17.21",
"socket.io-client": "^4.7.2",
"sonner": "^1.7.1",
"swr": "^2.2.2",
"uuid": "^9.0.0"
},
"pnpm": {
"overrides": {
"vite@>=4.4.0 <4.4.12": ">=4.4.12",
"@adobe/css-tools@<4.3.2": ">=4.3.2",
"vite@>=4.0.0 <=4.5.1": ">=4.5.2",
"vite@>=4.0.0 <=4.5.2": ">=4.5.3",
"braces@<3.0.3": ">=3.0.3",
"ws@>=8.0.0 <8.17.1": ">=8.17.1",
"micromatch@<4.0.8": ">=4.0.8",
"vite@>=4.0.0 <4.5.4": ">=4.5.4",
"vite@>=4.0.0 <=4.5.3": ">=4.5.4",
"rollup@>=3.0.0 <3.29.5": ">=3.29.5",
"cross-spawn@>=7.0.0 <7.0.5": ">=7.0.5"
}
}
}
+3920
View File
File diff suppressed because it is too large Load Diff
+89
View File
@@ -0,0 +1,89 @@
import { useContext, useMemo } from 'react';
import { ChainlitAPI } from 'src/api';
import { ChainlitContext } from 'src/context';
import useSWR, { SWRConfig, SWRConfiguration } from 'swr';
import { useAuthState } from './auth/state';
const fetcher = async (client: ChainlitAPI, endpoint: string) => {
const res = await client.get(endpoint);
return res?.json();
};
const cloneClient = (client: ChainlitAPI): ChainlitAPI => {
// Shallow clone API client.
// TODO: Move me to core API.
// Create new client
const newClient = new ChainlitAPI('', 'webapp');
// Assign old properties to new client
Object.assign(newClient, client);
return newClient;
};
/**
* React hook for cached API data fetching using SWR (stale-while-revalidate).
* Optimized for GET requests with automatic caching and revalidation.
*
* Key features:
* - Automatic data caching and revalidation
* - Integration with React component lifecycle
* - Loading state management
* - Recoil state integration for global state
* - Memoized fetcher function to prevent unnecessary rerenders
*
* @param path - API endpoint path or null to disable the request
* @param config - Optional SWR configuration
* @returns SWR response object containing:
* - data: The fetched data
* - error: Any error that occurred
* - isValidating: Whether a request is in progress
* - mutate: Function to mutate the cached data
*
* @example
* const { data, error, isValidating } = useApi<UserData>('/user');
*/
function useApi<T>(
path?: string | null,
{ ...swrConfig }: SWRConfiguration = {}
) {
const client = useContext(ChainlitContext);
const { setUser } = useAuthState();
// Memoize the fetcher function to avoid recreating it on every render
const memoizedFetcher = useMemo(
() =>
([url]: [url: string]) => {
if (!swrConfig.onErrorRetry) {
swrConfig.onErrorRetry = (...args) => {
const [err] = args;
// Don't do automatic retry for 401 - it just means we're not logged in (yet).
if (err.status === 401) {
setUser(null);
return;
}
// Fall back to default behavior.
return SWRConfig.defaultValue.onErrorRetry(...args);
};
}
const useApiClient = cloneClient(client);
useApiClient.on401 = useApiClient.onError = undefined;
return fetcher(useApiClient, url);
},
[client]
);
// Use a stable key for useSWR
const swrKey = useMemo(() => {
return path ? [path] : null;
}, [path]);
return useSWR<T, Error>(swrKey, memoizedFetcher, swrConfig);
}
export { useApi, fetcher };
@@ -0,0 +1,20 @@
import { useEffect } from 'react';
import { IAuthConfig } from 'src/index';
import { useApi } from '../api';
import { useAuthState } from './state';
export const useAuthConfig = () => {
const { authConfig, setAuthConfig } = useAuthState();
const { data: authConfigData, isLoading } = useApi<IAuthConfig>(
authConfig ? null : '/auth/config'
);
useEffect(() => {
if (authConfigData) {
setAuthConfig(authConfigData);
}
}, [authConfigData, setAuthConfig]);
return { authConfig, isLoading };
};
@@ -0,0 +1,36 @@
import { IAuthConfig, IUser } from 'src/types';
import { useAuthConfig } from './config';
import { useSessionManagement } from './sessionManagement';
import { useUserManagement } from './userManagement';
export const useAuth = () => {
const { authConfig } = useAuthConfig();
const { logout } = useSessionManagement();
const { user, setUserFromAPI } = useUserManagement();
const isReady =
!!authConfig && (!authConfig.requireLogin || user !== undefined);
if (authConfig && !authConfig.requireLogin) {
return {
data: authConfig,
user: null,
isReady,
isAuthenticated: true,
logout: () => Promise.resolve(),
setUserFromAPI: () => Promise.resolve()
};
}
return {
data: authConfig,
user,
isReady,
isAuthenticated: !!user,
logout,
setUserFromAPI
};
};
export type { IAuthConfig, IUser };
@@ -0,0 +1,21 @@
import { useContext } from 'react';
import { ChainlitContext } from 'src/index';
import { useAuthState } from './state';
export const useSessionManagement = () => {
const apiClient = useContext(ChainlitContext);
const { setUser, setThreadHistory } = useAuthState();
const logout = async (reload = false): Promise<void> => {
await apiClient.logout();
setUser(undefined);
setThreadHistory(undefined);
if (reload) {
window.location.reload();
}
};
return { logout };
};
@@ -0,0 +1,16 @@
import { useRecoilState, useSetRecoilState } from 'recoil';
import { authState, threadHistoryState, userState } from 'src/state';
export const useAuthState = () => {
const [authConfig, setAuthConfig] = useRecoilState(authState);
const [user, setUser] = useRecoilState(userState);
const setThreadHistory = useSetRecoilState(threadHistoryState);
return {
authConfig,
setAuthConfig,
user,
setUser,
setThreadHistory
};
};
@@ -0,0 +1,19 @@
import { IAuthConfig, IUser } from 'src/types';
export interface JWTPayload extends IUser {
exp: number;
}
export interface AuthState {
data: IAuthConfig | undefined;
user: IUser | null;
isAuthenticated: boolean;
isReady: boolean;
}
export interface AuthActions {
logout: (reload?: boolean) => Promise<void>;
setUserFromAPI: () => Promise<void>;
}
export type IUseAuth = AuthState & AuthActions;
@@ -0,0 +1,29 @@
import { useEffect } from 'react';
import { IUser } from 'src/types';
import { useApi } from '../api';
import { useAuthState } from './state';
export const useUserManagement = () => {
const { user, setUser } = useAuthState();
const {
data: userData,
error,
mutate: setUserFromAPI
} = useApi<IUser>('/user');
useEffect(() => {
if (userData) {
setUser(userData);
}
}, [userData, setUser]);
useEffect(() => {
if (error) {
setUser(null);
}
}, [error]);
return { user, setUserFromAPI };
};
+389
View File
@@ -0,0 +1,389 @@
import { IElement, IThread, IUser } from 'src/types';
import { IAction } from 'src/types/action';
import { IFeedback } from 'src/types/feedback';
export * from './hooks/auth';
export * from './hooks/api';
export interface IThreadFilters {
search?: string;
feedback?: number;
}
export interface IPageInfo {
hasNextPage: boolean;
endCursor?: string;
}
export interface IPagination {
first: number;
cursor?: string | number;
}
export class ClientError extends Error {
status: number;
detail?: string;
constructor(message: string, status: number, detail?: string) {
super(message);
this.status = status;
this.detail = detail;
}
toString() {
if (this.detail) {
return `${this.message}: ${this.detail}`;
} else {
return this.message;
}
}
}
type Payload = FormData | any;
export class APIBase {
constructor(
public httpEndpoint: string,
public type: 'webapp' | 'copilot' | 'teams' | 'slack' | 'discord',
public additionalQueryParams?: Record<string, string>,
public on401?: () => void,
public onError?: (error: ClientError) => void
) {}
buildEndpoint(path: string) {
let fullUrl = `${this.httpEndpoint}${path}`;
if (this.httpEndpoint.endsWith('/')) {
// remove trailing slash on httpEndpoint
fullUrl = `${this.httpEndpoint.slice(0, -1)}${path}`;
}
const url = new URL(fullUrl);
// Add additionalQueryParams for all API calls
if (this.additionalQueryParams) {
const params = new URLSearchParams(this.additionalQueryParams);
const separator = url.search ? '&' : '?';
url.search = url.search + `${separator}${params.toString()}`;
}
return url.toString();
}
private async getDetailFromErrorResponse(
res: Response
): Promise<string | undefined> {
try {
const body = await res.json();
return body?.detail;
} catch (error: any) {
console.error('Unable to parse error response', error);
}
return undefined;
}
private handleRequestError(error: any) {
if (error instanceof ClientError) {
if (error.status === 401 && this.on401) {
this.on401();
}
if (this.onError) {
this.onError(error);
}
}
console.error(error);
}
/**
* Low-level HTTP request handler for direct API interactions.
* Provides full control over HTTP methods, request configuration, and error handling.
*
* Key features:
* - Supports all HTTP methods (GET, POST, PUT, PATCH, DELETE)
* - Handles both FormData and JSON payloads
* - Manages authentication headers
* - Custom error handling with ClientError class
* - Support for request cancellation via AbortSignal
*
* @param method - HTTP method to use (GET, POST, etc.)
* @param path - API endpoint path
* @param data - Optional request payload (FormData or JSON-serializable data)
* @param signal - Optional AbortSignal for request cancellation
* @returns Promise<Response>
* @throws ClientError for HTTP errors, including 401 unauthorized
*/
async fetch(
method: string,
path: string,
data?: Payload,
signal?: AbortSignal,
headers: { Authorization?: string; 'Content-Type'?: string } = {}
): Promise<Response> {
try {
let body;
if (data instanceof FormData) {
body = data;
} else {
headers['Content-Type'] = 'application/json';
body = data ? JSON.stringify(data) : null;
}
const res = await fetch(this.buildEndpoint(path), {
method,
credentials: 'include',
headers,
signal,
body
});
if (!res.ok) {
const detail = await this.getDetailFromErrorResponse(res);
throw new ClientError(res.statusText, res.status, detail);
}
return res;
} catch (error: any) {
this.handleRequestError(error);
throw error;
}
}
async get(endpoint: string) {
return await this.fetch('GET', endpoint);
}
async post(endpoint: string, data: Payload, signal?: AbortSignal) {
return await this.fetch('POST', endpoint, data, signal);
}
async put(endpoint: string, data: Payload) {
return await this.fetch('PUT', endpoint, data);
}
async patch(endpoint: string, data: Payload) {
return await this.fetch('PATCH', endpoint, data);
}
async delete(endpoint: string, data: Payload) {
return await this.fetch('DELETE', endpoint, data);
}
}
export class ChainlitAPI extends APIBase {
async headerAuth() {
const res = await this.post(`/auth/header`, {});
return res.json();
}
async jwtAuth(token: string) {
const res = await this.fetch('POST', '/auth/jwt', undefined, undefined, {
Authorization: `Bearer ${token}`
});
return res.json();
}
async stickyCookie(sessionId: string) {
const res = await this.fetch('POST', '/set-session-cookie', {
session_id: sessionId
});
return res.json();
}
async passwordAuth(data: FormData) {
const res = await this.post(`/login`, data);
return res.json();
}
async getUser(): Promise<IUser> {
const res = await this.get(`/user`);
return res.json();
}
async logout() {
const res = await this.post(`/logout`, {});
return res.json();
}
async setFeedback(
feedback: IFeedback,
sessionId: string
): Promise<{ success: boolean; feedbackId: string }> {
const res = await this.put(`/feedback`, { feedback, sessionId });
return res.json();
}
async deleteFeedback(feedbackId: string): Promise<{ success: boolean }> {
const res = await this.delete(`/feedback`, { feedbackId });
return res.json();
}
async listThreads(
pagination: IPagination,
filter: IThreadFilters
): Promise<{
pageInfo: IPageInfo;
data: IThread[];
}> {
const res = await this.post(`/project/threads`, { pagination, filter });
return res.json();
}
async renameThread(threadId: string, name: string) {
const res = await this.put(`/project/thread`, { threadId, name });
return res.json();
}
async deleteThread(threadId: string) {
const res = await this.delete(`/project/thread`, { threadId });
return res.json();
}
uploadFile(
file: File,
onProgress: (progress: number) => void,
sessionId: string,
parentId?: string
) {
const xhr = new XMLHttpRequest();
xhr.withCredentials = true;
const promise = new Promise<{ id: string }>((resolve, reject) => {
const formData = new FormData();
formData.append('file', file);
const ask_parent_id = parentId ? `&ask_parent_id=${parentId}` : '';
xhr.open(
'POST',
this.buildEndpoint(
`/project/file?session_id=${sessionId}${ask_parent_id}`
),
true
);
// Track the progress of the upload
xhr.upload.onprogress = function (event) {
if (event.lengthComputable) {
const percentage = (event.loaded / event.total) * 100;
onProgress(percentage);
}
};
xhr.onload = function () {
if (xhr.status === 200) {
const response = JSON.parse(xhr.responseText);
resolve(response);
return;
}
const contentType = xhr.getResponseHeader('Content-Type');
if (contentType && contentType.includes('application/json')) {
const response = JSON.parse(xhr.responseText);
reject(response.detail);
} else {
reject('Upload failed');
}
};
xhr.onerror = function () {
reject('Upload error');
};
xhr.send(formData);
});
return { xhr, promise };
}
async callAction(action: IAction, sessionId: string) {
const res = await this.post(`/project/action`, { sessionId, action });
return res.json();
}
async updateElement(element: IElement, sessionId: string) {
const res = await this.put(`/project/element`, { sessionId, element });
return res.json();
}
async deleteElement(element: IElement, sessionId: string) {
const res = await this.delete(`/project/element`, { sessionId, element });
return res.json();
}
async connectStdioMCP(sessionId: string, name: string, fullCommand: string) {
const res = await this.post(`/mcp`, {
sessionId,
name,
fullCommand,
clientType: 'stdio'
});
return res.json();
}
async connectSseMCP(
sessionId: string,
name: string,
url: string,
headers?: Record<string, string>
) {
const res = await this.post(`/mcp`, {
sessionId,
name,
url,
...(headers ? { headers } : {}),
clientType: 'sse'
});
return res.json();
}
async connectStreamableHttpMCP(
sessionId: string,
name: string,
url: string,
headers?: Record<string, string>
) {
const res = await this.post(`/mcp`, {
sessionId,
name,
url,
...(headers ? { headers } : {}),
clientType: 'streamable-http'
});
return res.json();
}
async disconnectMcp(sessionId: string, name: string) {
const res = await this.delete(`/mcp`, { sessionId, name });
return res.json();
}
getElementUrl(id: string, sessionId: string) {
const queryParams = `?session_id=${sessionId}`;
return this.buildEndpoint(`/project/file/${id}${queryParams}`);
}
getLogoEndpoint(theme: string, configuredLogoUrl?: string) {
if (configuredLogoUrl) return configuredLogoUrl;
return this.buildEndpoint(`/logo?theme=${theme}`);
}
getOAuthEndpoint(provider: string) {
return this.buildEndpoint(`/auth/oauth/${provider}`);
}
async shareThread(
threadId: string,
isShared: boolean
): Promise<{ success: boolean }> {
const res = await this.put(`/project/thread/share`, {
threadId,
isShared
});
return res.json();
}
}
+11
View File
@@ -0,0 +1,11 @@
import { createContext } from 'react';
import { ChainlitAPI } from './api';
const defaultChainlitContext = undefined;
const ChainlitContext = createContext<ChainlitAPI>(
new ChainlitAPI('http://localhost:8000', 'webapp')
);
export { ChainlitContext, defaultChainlitContext };
+15
View File
@@ -0,0 +1,15 @@
export * from './useChatData';
export * from './useChatInteract';
export * from './useChatMessages';
export * from './useChatSession';
export * from './useAudio';
export * from './useConfig';
export * from './api';
export * from './types';
export * from './context';
export * from './state';
export * from './utils/message';
export { Socket } from 'socket.io-client';
export { WavRenderer } from './wavtools/wav_renderer';
+278
View File
@@ -0,0 +1,278 @@
import { isEqual } from 'lodash';
import { AtomEffect, DefaultValue, atom, selector } from 'recoil';
import { Socket } from 'socket.io-client';
import { v4 as uuidv4 } from 'uuid';
import { ICommand } from './types/command';
import { IMode } from './types/mode';
import {
IAction,
IAsk,
IAuthConfig,
ICallFn,
IChainlitConfig,
IMcp,
IMessageElement,
IStep,
ITasklistElement,
IUser,
ThreadHistory
} from './types';
import { groupByDate } from './utils/group';
import { WavRecorder, WavStreamPlayer } from './wavtools';
export interface ISession {
socket: Socket;
error?: boolean;
}
export const threadIdToResumeState = atom<string | undefined>({
key: 'ThreadIdToResume',
default: undefined
});
export const resumeThreadErrorState = atom<string | undefined>({
key: 'ResumeThreadErrorState',
default: undefined
});
export const chatProfileState = atom<string | undefined>({
key: 'ChatProfile',
default: undefined
});
const sessionIdAtom = atom<string>({
key: 'SessionId',
default: uuidv4()
});
export const sessionIdState = selector({
key: 'SessionIdSelector',
get: ({ get }) => get(sessionIdAtom),
set: ({ set }, newValue) =>
set(sessionIdAtom, newValue instanceof DefaultValue ? uuidv4() : newValue)
});
export const sessionState = atom<ISession | undefined>({
key: 'Session',
dangerouslyAllowMutability: true,
default: undefined
});
export const actionState = atom<IAction[]>({
key: 'Actions',
default: []
});
export const messagesState = atom<IStep[]>({
key: 'Messages',
dangerouslyAllowMutability: true,
default: []
});
export const commandsState = atom<ICommand[]>({
key: 'Commands',
default: []
});
export const modesState = atom<IMode[]>({
key: 'Modes',
default: []
});
export const tokenCountState = atom<number>({
key: 'TokenCount',
default: 0
});
export const loadingState = atom<boolean>({
key: 'Loading',
default: false
});
export const askUserState = atom<IAsk | undefined>({
key: 'AskUser',
default: undefined
});
export const wavRecorderState = atom({
key: 'WavRecorder',
dangerouslyAllowMutability: true,
default: new WavRecorder()
});
export const wavStreamPlayerState = atom({
key: 'WavStreamPlayer',
dangerouslyAllowMutability: true,
default: new WavStreamPlayer()
});
export const audioConnectionState = atom<'connecting' | 'on' | 'off'>({
key: 'AudioConnection',
default: 'off'
});
export const isAiSpeakingState = atom({
key: 'isAiSpeaking',
default: false
});
export const callFnState = atom<ICallFn | undefined>({
key: 'CallFn',
default: undefined
});
export const chatSettingsInputsState = atom<any>({
key: 'ChatSettings',
default: []
});
export const chatSettingsDefaultValueSelector = selector({
key: 'ChatSettingsValue/Default',
get: ({ get }) => {
const chatSettings = get(chatSettingsInputsState);
const collectInitialValues = (
inputs: any[],
acc: Record<string, any>
): Record<string, any> => {
if (!Array.isArray(inputs)) {
return acc;
}
inputs.forEach((input) => {
if (!input) {
return;
}
if (Array.isArray(input?.inputs) && input.inputs.length > 0) {
// Handle tabs
collectInitialValues(input.inputs, acc);
} else if (input?.id !== undefined) {
acc[input.id] = input.initial;
}
});
return acc;
};
return collectInitialValues(chatSettings, {});
}
});
export const chatSettingsValueState = atom<Record<string, any>>({
key: 'ChatSettingsValue',
default: chatSettingsDefaultValueSelector
});
export const elementState = atom<IMessageElement[]>({
key: 'DisplayElements',
default: []
});
export const tasklistState = atom<ITasklistElement[]>({
key: 'TasklistElements',
default: []
});
export const firstUserInteraction = atom<string | undefined>({
key: 'FirstUserInteraction',
default: undefined
});
export const userState = atom<IUser | undefined | null>({
key: 'User',
default: undefined
});
export const configState = atom<IChainlitConfig | undefined>({
key: 'ChainlitConfig',
default: undefined
});
export const authState = atom<IAuthConfig | undefined>({
key: 'AuthConfig',
default: undefined
});
export const threadHistoryState = atom<ThreadHistory | undefined>({
key: 'ThreadHistory',
default: {
threads: undefined,
currentThreadId: undefined,
timeGroupedThreads: undefined,
pageInfo: undefined
},
effects: [
({ setSelf, onSet }: { setSelf: any; onSet: any }) => {
onSet(
(
newValue: ThreadHistory | undefined,
oldValue: ThreadHistory | undefined
) => {
let timeGroupedThreads = newValue?.timeGroupedThreads;
if (
newValue?.threads &&
!isEqual(newValue.threads, oldValue?.timeGroupedThreads)
) {
timeGroupedThreads = groupByDate(newValue.threads);
}
setSelf({
...newValue,
timeGroupedThreads
});
}
);
}
]
});
export const sideViewState = atom<
{ title: string; elements: IMessageElement[]; key?: string } | undefined
>({
key: 'SideView',
default: undefined
});
export const currentThreadIdState = atom<string | undefined>({
key: 'CurrentThreadId',
default: undefined
});
const localStorageEffect =
<T>(key: string): AtomEffect<T> =>
({ setSelf, onSet }) => {
// When the atom is first initialized, try to get its value from localStorage
const savedValue = localStorage.getItem(key);
if (savedValue != null) {
try {
setSelf(JSON.parse(savedValue));
} catch (error) {
console.error(
`Error parsing localStorage value for key "${key}":`,
error
);
}
}
// Subscribe to state changes and update localStorage
onSet((newValue, _, isReset) => {
if (isReset) {
localStorage.removeItem(key);
} else {
localStorage.setItem(key, JSON.stringify(newValue));
}
});
};
export const mcpState = atom<IMcp[]>({
key: 'Mcp',
default: [],
effects: [localStorageEffect<IMcp[]>('mcp_storage_key')]
});
export const favoriteMessagesState = atom<IStep[]>({
key: 'favoriteMessagesState',
default: []
});
+16
View File
@@ -0,0 +1,16 @@
export interface IAction {
label: string;
forId: string;
id: string;
payload: Record<string, unknown>;
name: string;
onClick: () => void;
tooltip: string;
icon?: string;
}
export interface ICallFn {
callback: (payload: Record<string, any>) => void;
name: string;
args: Record<string, any>;
}
+5
View File
@@ -0,0 +1,5 @@
export interface OutputAudioChunk {
track: string;
mimeType: string;
data: Int16Array;
}
+8
View File
@@ -0,0 +1,8 @@
export interface ICommand {
id: string;
icon: string;
description: string;
button?: boolean;
persistent?: boolean;
selected?: boolean;
}
+108
View File
@@ -0,0 +1,108 @@
export interface IStarter {
label: string;
message: string;
icon?: string;
command?: string;
}
export interface IStarterCategory {
label: string;
icon?: string;
starters: IStarter[];
}
export interface ChatProfile {
default: boolean;
icon?: string;
name: string;
display_name?: string;
markdown_description: string;
starters?: IStarter[];
}
export interface IAudioConfig {
enabled: boolean;
sample_rate: number;
}
export interface IAuthConfig {
requireLogin: boolean;
passwordAuth: boolean;
headerAuth: boolean;
oauthProviders: string[];
default_theme?: 'light' | 'dark';
ui?: IChainlitConfig['ui'];
}
export interface IChainlitConfig {
markdown?: string;
ui: {
name: string;
description?: string;
default_theme?: 'light' | 'dark';
layout?: 'default' | 'wide';
default_sidebar_state?: 'open' | 'closed' | 'hidden';
chat_settings_location?: 'message_composer' | 'sidebar';
default_chat_settings_open?: boolean;
confirm_new_chat?: boolean;
cot: 'hidden' | 'tool_call' | 'full';
github?: string;
custom_css?: string;
custom_js?: string;
custom_font?: string;
alert_style?: 'classic' | 'modern';
login_page_image?: string;
login_page_image_filter?: string;
login_page_image_dark_filter?: string;
custom_meta_image_url?: string;
logo_file_url?: string;
default_avatar_file_url?: string;
avatar_size?: number;
header_links?: {
name: string;
display_name: string;
icon_url: string;
url: string;
target?: '_blank' | '_self' | '_parent' | '_top';
}[];
};
features: {
spontaneous_file_upload?: {
enabled?: boolean;
max_size_mb?: number;
max_files?: number;
accept?: string[] | Record<string, string[]>;
};
audio: IAudioConfig;
unsafe_allow_html?: boolean;
user_message_autoscroll?: boolean;
assistant_message_autoscroll?: boolean;
latex?: boolean;
user_message_markdown?: boolean;
edit_message?: boolean;
favorites?: boolean;
mcp?: {
enabled?: boolean;
sse?: {
enabled?: boolean;
};
streamable_http?: {
enabled?: boolean;
};
stdio?: {
enabled?: boolean;
};
};
};
debugUrl?: string;
userEnv: string[];
maskUserEnv?: boolean;
dataPersistence: boolean;
threadResumable: boolean;
threadSharing?: boolean;
chatProfiles: ChatProfile[];
starters?: IStarter[];
starterCategories?: IStarterCategory[];
translation: object;
}
+81
View File
@@ -0,0 +1,81 @@
export type IElement =
| IImageElement
| ITextElement
| IPdfElement
| ITasklistElement
| IAudioElement
| IVideoElement
| IFileElement
| IPlotlyElement
| IDataframeElement
| ICustomElement;
export type IMessageElement =
| IImageElement
| ITextElement
| IPdfElement
| IAudioElement
| IVideoElement
| IFileElement
| IPlotlyElement
| IDataframeElement
| ICustomElement;
export type ElementType = IElement['type'];
export type IElementSize = 'small' | 'medium' | 'large';
interface TElement<T> {
id: string;
type: T;
threadId?: string;
forId: string;
mime?: string;
url?: string;
chainlitKey?: string;
}
interface TMessageElement<T> extends TElement<T> {
name: string;
display: 'inline' | 'side' | 'page';
}
export interface IImageElement extends TMessageElement<'image'> {
size?: IElementSize;
}
export interface ITextElement extends TMessageElement<'text'> {
language?: string;
}
export interface IPdfElement extends TMessageElement<'pdf'> {
page?: number;
}
export interface IAudioElement extends TMessageElement<'audio'> {
autoPlay?: boolean;
}
export interface IVideoElement extends TMessageElement<'video'> {
size?: IElementSize;
/**
* Override settings for each type of player in ReactPlayer
* https://github.com/cookpete/react-player?tab=readme-ov-file#config-prop
* @type {object}
*/
playerConfig?: object;
}
export interface IFileElement extends TMessageElement<'file'> {
type: 'file';
}
export type IPlotlyElement = TMessageElement<'plotly'>;
export type ITasklistElement = TElement<'tasklist'>;
export type IDataframeElement = TMessageElement<'dataframe'>;
export interface ICustomElement extends TMessageElement<'custom'> {
props: Record<string, unknown>;
}
+7
View File
@@ -0,0 +1,7 @@
export interface IFeedback {
id?: string;
forId?: string;
threadId?: string;
comment?: string;
value: number;
}
+35
View File
@@ -0,0 +1,35 @@
import { IAction } from './action';
import { IStep } from './step';
export interface IAskElementResponse {
submitted: boolean;
[key: string]: unknown;
}
export interface FileSpec {
accept?: string[] | Record<string, string[]>;
max_size_mb?: number;
max_files?: number;
}
export interface ActionSpec {
keys?: string[];
}
export interface IFileRef {
id: string;
}
export interface IAsk {
callback: (
payload: IStep | IFileRef[] | IAction | IAskElementResponse
) => void;
spec: {
type: 'text' | 'file' | 'action' | 'element';
step_id: string;
timeout: number;
element_id?: string;
} & FileSpec &
ActionSpec;
parentId?: string;
}
+15
View File
@@ -0,0 +1,15 @@
import { IThread } from 'src/types';
import { IPageInfo } from '..';
export type UserInput = {
content: string;
createdAt: number;
};
export type ThreadHistory = {
threads?: IThread[];
currentThreadId?: string;
timeGroupedThreads?: { [key: string]: IThread[] };
pageInfo?: IPageInfo;
};
+12
View File
@@ -0,0 +1,12 @@
export * from './action';
export * from './element';
export * from './command';
export * from './mode';
export * from './file';
export * from './feedback';
export * from './step';
export * from './user';
export * from './thread';
export * from './history';
export * from './config';
export * from './mcp';
+10
View File
@@ -0,0 +1,10 @@
export interface IMcp {
name: string;
tools: [{ name: string }];
status: 'connected' | 'connecting' | 'failed';
clientType: 'sse' | 'stdio' | 'streamable-http';
command?: string;
url?: string;
/** Optional HTTP headers used when connecting (SSE or streamable-http) */
headers?: Record<string, string>;
}
+20
View File
@@ -0,0 +1,20 @@
/**
* Represents a single selectable option within a mode.
*/
export interface IModeOption {
id: string;
name: string;
description?: string;
icon?: string;
default?: boolean;
}
/**
* Represents a mode category containing multiple selectable options.
* Examples: Model selection, Reasoning Effort, Approach preference, etc.
*/
export interface IMode {
id: string;
name: string;
options: IModeOption[];
}
+40
View File
@@ -0,0 +1,40 @@
import { IFeedback } from './feedback';
type StepType =
| 'assistant_message'
| 'user_message'
| 'system_message'
| 'run'
| 'tool'
| 'llm'
| 'embedding'
| 'retrieval'
| 'rerank'
| 'undefined';
export interface IStep {
id: string;
name: string;
type: StepType;
threadId?: string;
parentId?: string;
isError?: boolean;
command?: string;
modes?: Record<string, string>;
showInput?: boolean | string;
waitForAnswer?: boolean;
input?: string;
output: string;
createdAt: number | string;
start?: number | string;
end?: number | string;
feedback?: IFeedback;
language?: string;
defaultOpen?: boolean;
autoCollapse?: boolean;
streaming?: boolean;
steps?: IStep[];
metadata?: Record<string, any>;
//legacy
indent?: number;
}
+13
View File
@@ -0,0 +1,13 @@
import { IElement } from './element';
import { IStep } from './step';
export interface IThread {
id: string;
createdAt: number | string;
name?: string;
userId?: string;
userIdentifier?: string;
metadata?: Record<string, any>;
steps: IStep[];
elements?: IElement[];
}
+20
View File
@@ -0,0 +1,20 @@
export type AuthProvider =
| 'credentials'
| 'header'
| 'github'
| 'google'
| 'azure-ad'
| 'azure-ad-hybrid';
export interface IUserMetadata extends Record<string, any> {
tags?: string[];
image?: string;
provider?: AuthProvider;
}
export interface IUser {
id: string;
identifier: string;
display_name?: string;
metadata: IUserMetadata;
}
+43
View File
@@ -0,0 +1,43 @@
import { useCallback } from 'react';
import { useRecoilState, useRecoilValue } from 'recoil';
import {
audioConnectionState,
isAiSpeakingState,
wavRecorderState,
wavStreamPlayerState
} from './state';
import { useChatInteract } from './useChatInteract';
const useAudio = () => {
const [audioConnection, setAudioConnection] =
useRecoilState(audioConnectionState);
const wavRecorder = useRecoilValue(wavRecorderState);
const wavStreamPlayer = useRecoilValue(wavStreamPlayerState);
const isAiSpeaking = useRecoilValue(isAiSpeakingState);
const { startAudioStream, endAudioStream } = useChatInteract();
const startConversation = useCallback(async () => {
setAudioConnection('connecting');
await startAudioStream();
}, [startAudioStream]);
const endConversation = useCallback(async () => {
setAudioConnection('off');
await wavRecorder.end();
await wavStreamPlayer.interrupt();
await endAudioStream();
}, [endAudioStream, wavRecorder, wavStreamPlayer]);
return {
startConversation,
endConversation,
audioConnection,
isAiSpeaking,
wavRecorder,
wavStreamPlayer
};
};
export { useAudio };
+61
View File
@@ -0,0 +1,61 @@
import { useRecoilValue } from 'recoil';
import {
actionState,
askUserState,
callFnState,
chatSettingsDefaultValueSelector,
chatSettingsInputsState,
chatSettingsValueState,
elementState,
loadingState,
sessionState,
tasklistState
} from './state';
export interface IToken {
id: number | string;
token: string;
isSequence: boolean;
isInput: boolean;
}
const useChatData = () => {
const loading = useRecoilValue(loadingState);
const elements = useRecoilValue(elementState);
const tasklists = useRecoilValue(tasklistState);
const actions = useRecoilValue(actionState);
const session = useRecoilValue(sessionState);
const askUser = useRecoilValue(askUserState);
const callFn = useRecoilValue(callFnState);
const chatSettingsInputs = useRecoilValue(chatSettingsInputsState);
const chatSettingsValue = useRecoilValue(chatSettingsValueState);
const chatSettingsDefaultValue = useRecoilValue(
chatSettingsDefaultValueSelector
);
const connected = session?.socket.connected && !session?.error;
const disabled =
!connected ||
loading ||
askUser?.spec.type === 'file' ||
askUser?.spec.type === 'action' ||
askUser?.spec.type === 'element';
return {
actions,
askUser,
callFn,
chatSettingsDefaultValue,
chatSettingsInputs,
chatSettingsValue,
connected,
disabled,
elements,
error: session?.error,
loading,
tasklists
};
};
export { useChatData };
+224
View File
@@ -0,0 +1,224 @@
import { useCallback, useContext } from 'react';
import { useRecoilValue, useResetRecoilState, useSetRecoilState } from 'recoil';
import {
actionState,
askUserState,
chatSettingsInputsState,
chatSettingsValueState,
currentThreadIdState,
elementState,
favoriteMessagesState,
firstUserInteraction,
loadingState,
messagesState,
sessionIdState,
sessionState,
sideViewState,
tasklistState,
threadIdToResumeState,
tokenCountState
} from 'src/state';
import { IFileRef, IStep } from 'src/types';
import { addMessage } from 'src/utils/message';
import { v4 as uuidv4 } from 'uuid';
import { ChainlitContext } from './context';
type PartialBy<T, K extends keyof T> = Omit<T, K> & Partial<Pick<T, K>>;
const useChatInteract = () => {
const client = useContext(ChainlitContext);
const session = useRecoilValue(sessionState);
const askUser = useRecoilValue(askUserState);
const sessionId = useRecoilValue(sessionIdState);
const resetChatSettings = useResetRecoilState(chatSettingsInputsState);
const resetSessionId = useResetRecoilState(sessionIdState);
const resetChatSettingsValue = useResetRecoilState(chatSettingsValueState);
const setFirstUserInteraction = useSetRecoilState(firstUserInteraction);
const setLoading = useSetRecoilState(loadingState);
const setMessages = useSetRecoilState(messagesState);
const setElements = useSetRecoilState(elementState);
const setTasklists = useSetRecoilState(tasklistState);
const setActions = useSetRecoilState(actionState);
const setTokenCount = useSetRecoilState(tokenCountState);
const setIdToResume = useSetRecoilState(threadIdToResumeState);
const setSideView = useSetRecoilState(sideViewState);
const setCurrentThreadId = useSetRecoilState(currentThreadIdState);
const setFavoriteMessages = useSetRecoilState(favoriteMessagesState);
const clear = useCallback(() => {
session?.socket.emit('clear_session');
session?.socket.disconnect();
setIdToResume(undefined);
resetSessionId();
setFirstUserInteraction(undefined);
setMessages([]);
setElements([]);
setTasklists([]);
setActions([]);
setTokenCount(0);
resetChatSettings();
resetChatSettingsValue();
setSideView(undefined);
setCurrentThreadId(undefined);
}, [session]);
const sendMessage = useCallback(
(
message: PartialBy<IStep, 'createdAt' | 'id'>,
fileReferences: IFileRef[] = []
) => {
if (!message.id) {
message.id = uuidv4();
}
if (!message.createdAt) {
message.createdAt = new Date().toISOString();
}
setMessages((oldMessages) => addMessage(oldMessages, message as IStep));
session?.socket.emit('client_message', { message, fileReferences });
},
[session?.socket]
);
const editMessage = useCallback(
(message: IStep) => {
session?.socket.emit('edit_message', { message });
},
[session?.socket]
);
const toggleMessageFavorite = useCallback(
(message: IStep) => {
const favorite = !(message.metadata?.favorite ?? false);
const updatedMetadata = {
...(message.metadata || {}),
favorite
};
setMessages((oldMessages) =>
oldMessages.map((item) =>
item.id === message.id
? { ...item, metadata: { ...(item.metadata || {}), favorite } }
: item
)
);
const nextMessage: IStep = {
...message,
metadata: updatedMetadata
};
setFavoriteMessages((oldFavorites) => {
if (favorite) {
const filtered = oldFavorites.filter(
(step) => step.id !== message.id
);
return [nextMessage, ...filtered];
}
return oldFavorites.filter((step) => step.id !== message.id);
});
session?.socket.emit('message_favorite', { message: nextMessage });
},
[session?.socket, setFavoriteMessages, setMessages]
);
const windowMessage = useCallback(
(data: any) => {
session?.socket.emit('window_message', data);
},
[session?.socket]
);
const startAudioStream = useCallback(() => {
session?.socket.emit('audio_start');
}, [session?.socket]);
const sendAudioChunk = useCallback(
(
isStart: boolean,
mimeType: string,
elapsedTime: number,
data: Int16Array
) => {
session?.socket.emit('audio_chunk', {
isStart,
mimeType,
elapsedTime,
data
});
},
[session?.socket]
);
const endAudioStream = useCallback(() => {
session?.socket.emit('audio_end');
}, [session?.socket]);
const replyMessage = useCallback(
(message: IStep) => {
if (askUser) {
if (askUser.parentId) message.parentId = askUser.parentId;
setMessages((oldMessages) => addMessage(oldMessages, message));
askUser.callback(message);
}
},
[askUser]
);
const updateChatSettings = useCallback(
(values: object) => {
session?.socket.emit('chat_settings_change', values);
},
[session?.socket]
);
const editChatSettings = useCallback(
(values: object) => {
session?.socket.emit('chat_settings_edit', values);
},
[session?.socket]
);
const stopTask = useCallback(() => {
setMessages((oldMessages) =>
oldMessages.map((m) => {
m.streaming = false;
return m;
})
);
setLoading(false);
session?.socket.emit('stop');
}, [session?.socket]);
const uploadFile = useCallback(
(file: File, onProgress: (progress: number) => void, parentId?: string) => {
return client.uploadFile(file, onProgress, sessionId, parentId);
},
[sessionId]
);
return {
uploadFile,
clear,
replyMessage,
sendMessage,
editMessage,
windowMessage,
startAudioStream,
sendAudioChunk,
endAudioStream,
stopTask,
setIdToResume,
updateChatSettings,
editChatSettings,
toggleMessageFavorite
};
};
export { useChatInteract };
+21
View File
@@ -0,0 +1,21 @@
import { useRecoilValue } from 'recoil';
import {
currentThreadIdState,
firstUserInteraction,
messagesState
} from './state';
const useChatMessages = () => {
const messages = useRecoilValue(messagesState);
const firstInteraction = useRecoilValue(firstUserInteraction);
const threadId = useRecoilValue(currentThreadIdState);
return {
threadId,
messages,
firstInteraction
};
};
export { useChatMessages };
+519
View File
@@ -0,0 +1,519 @@
import { debounce } from 'lodash';
import { useCallback, useContext, useEffect } from 'react';
import {
useRecoilState,
useRecoilValue,
useResetRecoilState,
useSetRecoilState
} from 'recoil';
import io from 'socket.io-client';
import { toast } from 'sonner';
import {
actionState,
askUserState,
audioConnectionState,
callFnState,
chatProfileState,
chatSettingsInputsState,
chatSettingsValueState,
commandsState,
currentThreadIdState,
elementState,
favoriteMessagesState,
firstUserInteraction,
isAiSpeakingState,
loadingState,
mcpState,
messagesState,
modesState,
resumeThreadErrorState,
sessionIdState,
sessionState,
sideViewState,
tasklistState,
threadIdToResumeState,
tokenCountState,
wavRecorderState,
wavStreamPlayerState
} from 'src/state';
import {
IAction,
ICommand,
IElement,
IMessageElement,
IMode,
IStep,
ITasklistElement,
IThread
} from 'src/types';
import {
addMessage,
deleteMessageById,
updateMessageById,
updateMessageContentById
} from 'src/utils/message';
import { OutputAudioChunk } from './types/audio';
import { ChainlitContext } from './context';
import type { IToken } from './useChatData';
const useChatSession = () => {
const client = useContext(ChainlitContext);
const sessionId = useRecoilValue(sessionIdState);
const [session, setSession] = useRecoilState(sessionState);
const setIsAiSpeaking = useSetRecoilState(isAiSpeakingState);
const setAudioConnection = useSetRecoilState(audioConnectionState);
const resetChatSettingsValue = useResetRecoilState(chatSettingsValueState);
const setChatSettingsValue = useSetRecoilState(chatSettingsValueState);
const setFirstUserInteraction = useSetRecoilState(firstUserInteraction);
const setLoading = useSetRecoilState(loadingState);
const setMcps = useSetRecoilState(mcpState);
const wavStreamPlayer = useRecoilValue(wavStreamPlayerState);
const wavRecorder = useRecoilValue(wavRecorderState);
const setMessages = useSetRecoilState(messagesState);
const setAskUser = useSetRecoilState(askUserState);
const setCallFn = useSetRecoilState(callFnState);
const setCommands = useSetRecoilState(commandsState);
const setModes = useSetRecoilState(modesState);
const setSideView = useSetRecoilState(sideViewState);
const setElements = useSetRecoilState(elementState);
const setTasklists = useSetRecoilState(tasklistState);
const setActions = useSetRecoilState(actionState);
const setChatSettingsInputs = useSetRecoilState(chatSettingsInputsState);
const setTokenCount = useSetRecoilState(tokenCountState);
const [chatProfile, setChatProfile] = useRecoilState(chatProfileState);
const idToResume = useRecoilValue(threadIdToResumeState);
const setThreadResumeError = useSetRecoilState(resumeThreadErrorState);
const setFavoriteMessages = useSetRecoilState(favoriteMessagesState);
const [currentThreadId, setCurrentThreadId] =
useRecoilState(currentThreadIdState);
// Use currentThreadId as thread id in websocket header
useEffect(() => {
if (session?.socket) {
session.socket.auth['threadId'] = currentThreadId || '';
}
}, [currentThreadId]);
const _connect = useCallback(
async ({
transports,
userEnv
}: {
transports?: string[];
userEnv: Record<string, string>;
}) => {
const { protocol, host, pathname } = new URL(client.httpEndpoint);
const uri = `${protocol}//${host}`;
const path =
pathname && pathname !== '/'
? `${pathname}/ws/socket.io`
: '/ws/socket.io';
try {
await client.stickyCookie(sessionId);
} catch (err) {
console.error(`Failed to set sticky session cookie: ${err}`);
}
const socket = io(uri, {
path,
withCredentials: true,
transports,
auth: {
clientType: client.type,
sessionId,
threadId: idToResume || '',
userEnv: JSON.stringify(userEnv),
chatProfile: chatProfile ? encodeURIComponent(chatProfile) : ''
}
});
setSession((old) => {
old?.socket?.removeAllListeners();
old?.socket?.close();
return {
socket
};
});
socket.on('connect', () => {
socket.emit('connection_successful');
setSession((s) => ({ ...s!, error: false }));
socket.emit('fetch_favorites');
setMcps((prev) =>
prev.map((mcp) => {
let promise;
if (mcp.clientType === 'sse') {
promise = client.connectSseMCP(sessionId, mcp.name, mcp.url!);
} else if (mcp.clientType === 'streamable-http') {
promise = client.connectStreamableHttpMCP(
sessionId,
mcp.name,
mcp.url!,
mcp.headers || {}
);
} else {
promise = client.connectStdioMCP(
sessionId,
mcp.name,
mcp.command!
);
}
promise
.then(async ({ success, mcp }) => {
setMcps((prev) =>
prev.map((existingMcp) => {
if (existingMcp.name === mcp.name) {
return {
...existingMcp,
status: success ? 'connected' : 'failed',
tools: mcp ? mcp.tools : existingMcp.tools
};
}
return existingMcp;
})
);
})
.catch(() => {
setMcps((prev) =>
prev.map((existingMcp) => {
if (existingMcp.name === mcp.name) {
return {
...existingMcp,
status: 'failed'
};
}
return existingMcp;
})
);
});
return { ...mcp, status: 'connecting' };
})
);
});
socket.on('connect_error', (_) => {
setSession((s) => ({ ...s!, error: true }));
});
socket.on('task_start', () => {
setLoading(true);
});
socket.on('task_end', () => {
setLoading(false);
});
socket.on('reload', () => {
socket.emit('clear_session');
window.location.reload();
});
socket.on('audio_connection', async (state: 'on' | 'off') => {
if (state === 'on') {
let isFirstChunk = true;
const startTime = Date.now();
const mimeType = 'pcm16';
try {
await wavRecorder.begin();
await wavStreamPlayer.connect();
await wavRecorder.record(async (data) => {
const elapsedTime = Date.now() - startTime;
socket.emit('audio_chunk', {
isStart: isFirstChunk,
mimeType,
elapsedTime,
data: data.mono
});
isFirstChunk = false;
});
wavStreamPlayer.onStop = () => setIsAiSpeaking(false);
} catch {
try {
await wavRecorder.end();
} catch {
// ignored
}
await wavStreamPlayer.interrupt();
socket.emit('audio_end');
setAudioConnection('off');
return;
}
} else {
await wavRecorder.end();
await wavStreamPlayer.interrupt();
}
setAudioConnection(state);
});
socket.on('audio_chunk', (chunk: OutputAudioChunk) => {
wavStreamPlayer.add16BitPCM(chunk.data, chunk.track);
setIsAiSpeaking(true);
});
socket.on('audio_interrupt', () => {
wavStreamPlayer.interrupt();
});
socket.on('resume_thread', (thread: IThread) => {
const isReadOnlyView = Boolean(
(thread as any)?.metadata?.viewer_read_only
);
if (!isReadOnlyView && idToResume && thread.id !== idToResume) {
window.location.href = `/thread/${thread.id}`;
}
if (!isReadOnlyView && idToResume) {
setCurrentThreadId(thread.id);
}
let messages: IStep[] = [];
for (const step of thread.steps) {
messages = addMessage(messages, step);
}
if (thread.metadata?.chat_profile) {
setChatProfile(thread.metadata?.chat_profile);
}
if (thread.metadata?.chat_settings) {
setChatSettingsValue(thread.metadata?.chat_settings);
}
setMessages(messages);
const elements = thread.elements || [];
setTasklists(
(elements as ITasklistElement[]).filter((e) => e.type === 'tasklist')
);
setElements(
(elements as IMessageElement[]).filter(
(e) => ['avatar', 'tasklist'].indexOf(e.type) === -1
)
);
});
socket.on('resume_thread_error', (error?: string) => {
setThreadResumeError(error);
});
socket.on('new_message', (message: IStep) => {
setMessages((oldMessages) => addMessage(oldMessages, message));
});
socket.on(
'first_interaction',
(event: { interaction: string; thread_id: string }) => {
setFirstUserInteraction(event.interaction);
setCurrentThreadId(event.thread_id);
}
);
socket.on('update_message', (message: IStep) => {
setMessages((oldMessages) =>
updateMessageById(oldMessages, message.id, message)
);
});
socket.on('delete_message', (message: IStep) => {
setMessages((oldMessages) =>
deleteMessageById(oldMessages, message.id)
);
});
socket.on('stream_start', (message: IStep) => {
setMessages((oldMessages) => addMessage(oldMessages, message));
});
socket.on(
'stream_token',
({ id, token, isSequence, isInput }: IToken) => {
setMessages((oldMessages) =>
updateMessageContentById(
oldMessages,
id,
token,
isSequence,
isInput
)
);
}
);
socket.on('ask', ({ msg, spec }, callback) => {
setAskUser({ spec, callback, parentId: msg.parentId });
setMessages((oldMessages) => addMessage(oldMessages, msg));
setLoading(false);
});
socket.on('ask_timeout', () => {
setAskUser(undefined);
setLoading(false);
});
socket.on('clear_ask', () => {
setAskUser(undefined);
});
socket.on('call_fn', ({ name, args }, callback) => {
setCallFn({ name, args, callback });
});
socket.on('clear_call_fn', () => {
setCallFn(undefined);
});
socket.on('call_fn_timeout', () => {
setCallFn(undefined);
});
socket.on('chat_settings', (inputs: any) => {
setChatSettingsInputs(inputs);
resetChatSettingsValue();
});
socket.on('set_commands', (commands: ICommand[]) => {
setCommands(commands);
});
socket.on('set_modes', (modes: IMode[]) => {
setModes(modes);
});
socket.on('set_favorites', (steps: IStep[]) => {
setFavoriteMessages(steps);
});
socket.on('set_sidebar_title', (title: string) => {
setSideView((prev) => {
if (prev?.title === title) return prev;
return { title, elements: prev?.elements || [] };
});
});
socket.on(
'set_sidebar_elements',
({ elements, key }: { elements: IMessageElement[]; key?: string }) => {
if (!elements.length) {
setSideView(undefined);
} else {
elements.forEach((element) => {
if (!element.url && element.chainlitKey) {
element.url = client.getElementUrl(
element.chainlitKey,
sessionId
);
}
});
setSideView((prev) => {
if (prev?.key === key) return prev;
return { title: prev?.title || '', elements: elements, key };
});
}
}
);
socket.on('element', (element: IElement) => {
if (!element.url && element.chainlitKey) {
element.url = client.getElementUrl(element.chainlitKey, sessionId);
}
if (element.type === 'tasklist') {
setTasklists((old) => {
const index = old.findIndex((e) => e.id === element.id);
if (index === -1) {
return [...old, element];
} else {
return [...old.slice(0, index), element, ...old.slice(index + 1)];
}
});
} else {
setElements((old) => {
const index = old.findIndex((e) => e.id === element.id);
if (index === -1) {
return [...old, element];
} else {
return [...old.slice(0, index), element, ...old.slice(index + 1)];
}
});
}
});
socket.on('remove_element', (remove: { id: string }) => {
setElements((old) => {
return old.filter((e) => e.id !== remove.id);
});
setTasklists((old) => {
return old.filter((e) => e.id !== remove.id);
});
});
socket.on('action', (action: IAction) => {
setActions((old) => [...old, action]);
});
socket.on('remove_action', (action: IAction) => {
setActions((old) => {
const index = old.findIndex((a) => a.id === action.id);
if (index === -1) return old;
return [...old.slice(0, index), ...old.slice(index + 1)];
});
});
socket.on('token_usage', (count: number) => {
setTokenCount((old) => old + count);
});
socket.on('window_message', (data: any) => {
if (window.parent) {
window.parent.postMessage(data, '*');
}
});
socket.on('toast', (data: { message: string; type: string }) => {
if (!data.message) {
console.warn('No message received for toast.');
return;
}
switch (data.type) {
case 'info':
toast.info(data.message);
break;
case 'error':
toast.error(data.message);
break;
case 'success':
toast.success(data.message);
break;
case 'warning':
toast.warning(data.message);
break;
default:
toast(data.message);
break;
}
});
},
[setSession, sessionId, idToResume, chatProfile]
);
const connect = useCallback(debounce(_connect, 200), [_connect]);
const disconnect = useCallback(() => {
if (session?.socket) {
session.socket.removeAllListeners();
session.socket.close();
}
}, [session]);
return {
connect,
disconnect,
session,
sessionId,
chatProfile,
idToResume,
setChatProfile
};
};
export { useChatSession };
+43
View File
@@ -0,0 +1,43 @@
import { useEffect, useRef } from 'react';
import { useRecoilState, useRecoilValue } from 'recoil';
import { useApi, useAuth } from './api';
import { chatProfileState, configState } from './state';
import { IChainlitConfig } from './types';
const useConfig = () => {
const [config, setConfig] = useRecoilState(configState);
const { isAuthenticated } = useAuth();
const chatProfile = useRecoilValue(chatProfileState);
const language = navigator.language || 'en-US';
const prevChatProfileRef = useRef(chatProfile);
// Build the API URL with optional chat profile parameter
const apiUrl = isAuthenticated
? `/project/settings?language=${language}${chatProfile ? `&chat_profile=${encodeURIComponent(chatProfile)}` : ''}`
: null;
// Always fetch if we don't have config and we're authenticated
const shouldFetch = isAuthenticated && !config;
const { data, error, isLoading } = useApi<IChainlitConfig>(
shouldFetch ? apiUrl : null
);
useEffect(() => {
if (!data) return;
setConfig(data);
}, [data, setConfig]);
// Clear config when chat profile changes to force re-fetch
useEffect(() => {
if (prevChatProfileRef.current !== chatProfile) {
setConfig(undefined);
prevChatProfileRef.current = chatProfile;
}
}, [chatProfile, setConfig]);
return { config, error, isLoading, language };
};
export { useConfig };
+44
View File
@@ -0,0 +1,44 @@
import { IThread } from 'src/types';
export const groupByDate = (data: IThread[]) => {
const groupedData: { [key: string]: IThread[] } = {};
const today = new Date();
today.setHours(0, 0, 0, 0);
[...data]
.sort(
(a, b) =>
new Date(b.createdAt).getTime() - new Date(a.createdAt).getTime()
)
.forEach((item) => {
const threadDate = new Date(item.createdAt);
threadDate.setHours(0, 0, 0, 0);
const daysDiff = Math.floor(
(today.getTime() - threadDate.getTime()) / 86400000
);
const locale = navigator.language;
let category: string;
if (daysDiff === 0) {
category = 'Today';
} else if (daysDiff === 1) {
category = 'Yesterday';
} else if (daysDiff <= 7) {
category = 'Previous 7 days';
} else if (daysDiff <= 30) {
category = 'Previous 30 days';
} else {
category = threadDate.toLocaleString(locale, {
month: 'long',
year: 'numeric'
});
}
groupedData[category] ??= [];
groupedData[category].push(item);
});
return groupedData;
};
+245
View File
@@ -0,0 +1,245 @@
import { isEqual } from 'lodash';
import { IStep } from '..';
const nestMessages = (messages: IStep[]): IStep[] => {
let nestedMessages: IStep[] = [];
for (const message of messages) {
nestedMessages = addMessage(nestedMessages, message);
}
return nestedMessages;
};
const isLastMessage = (messages: IStep[], index: number) => {
if (messages.length - 1 === index) {
return true;
}
for (let i = index + 1; i < messages.length; i++) {
if (messages[i].streaming) {
continue;
} else {
return false;
}
}
return true;
};
// Nested messages utils
const addMessage = (messages: IStep[], message: IStep): IStep[] => {
if (hasMessageById(messages, message.id)) {
return updateMessageById(messages, message.id, message);
} else if ('parentId' in message && message.parentId) {
return addMessageToParent(messages, message.parentId, message);
} else if ('indent' in message && message.indent && message.indent > 0) {
return addIndentMessage(messages, message.indent, message);
} else {
return [...messages, message];
}
};
const addIndentMessage = (
messages: IStep[],
indent: number,
newMessage: IStep,
currentIndentation: number = 0
): IStep[] => {
if (messages.length === 0) {
return [newMessage];
}
const index = messages.length - 1;
const msg = messages[index];
const msgSteps = msg.steps || [];
if (currentIndentation + 1 === indent) {
const updatedMsg = {
...msg,
steps: [...msgSteps, newMessage]
};
const nextMessages = [...messages];
nextMessages[index] = updatedMsg;
return nextMessages;
} else {
const updatedSteps = addIndentMessage(
msgSteps,
indent,
newMessage,
currentIndentation + 1
);
if (updatedSteps === msgSteps) {
return messages;
}
const nextMessages = [...messages];
nextMessages[index] = { ...msg, steps: updatedSteps };
return nextMessages;
}
};
const addMessageToParent = (
messages: IStep[],
parentId: string,
newMessage: IStep
): IStep[] => {
let hasChanges = false;
const nextMessages = messages.map((msg) => {
if (isEqual(msg.id, parentId)) {
hasChanges = true;
return {
...msg,
steps: msg.steps ? [...msg.steps, newMessage] : [newMessage]
};
} else if (hasMessageById(messages, parentId) && msg.steps) {
const updatedSteps = addMessageToParent(msg.steps, parentId, newMessage);
if (updatedSteps !== msg.steps) {
hasChanges = true;
return { ...msg, steps: updatedSteps };
}
}
return msg;
});
return hasChanges ? nextMessages : messages;
};
const findMessageById = (
messages: IStep[],
messageId: string
): IStep | undefined => {
for (const message of messages) {
if (isEqual(message.id, messageId)) {
return message;
} else if (message.steps && message.steps.length > 0) {
const foundMessage = findMessageById(message.steps, messageId);
if (foundMessage) {
return foundMessage;
}
}
}
return undefined;
};
const hasMessageById = (messages: IStep[], messageId: string): boolean => {
return findMessageById(messages, messageId) !== undefined;
};
const updateMessageById = (
messages: IStep[],
messageId: string,
updatedMessage: IStep
): IStep[] => {
let hasChanges = false;
const nextMessages = messages.map((msg) => {
if (isEqual(msg.id, messageId)) {
hasChanges = true;
return { ...msg, ...updatedMessage };
} else if (msg.steps) {
const updatedSteps = updateMessageById(
msg.steps,
messageId,
updatedMessage
);
if (updatedSteps !== msg.steps) {
hasChanges = true;
return { ...msg, steps: updatedSteps };
}
}
return msg;
});
return hasChanges ? nextMessages : messages;
};
const deleteMessageById = (messages: IStep[], messageId: string): IStep[] => {
let hasChanges = false;
const nextMessages = messages.reduce((acc, msg) => {
if (msg.id === messageId) {
hasChanges = true;
return acc;
} else if (msg.steps) {
const updatedSteps = deleteMessageById(msg.steps, messageId);
if (updatedSteps !== msg.steps) {
hasChanges = true;
acc.push({ ...msg, steps: updatedSteps });
return acc;
}
}
acc.push(msg);
return acc;
}, [] as IStep[]);
return hasChanges ? nextMessages : messages;
};
const updateMessageContentById = (
messages: IStep[],
messageId: number | string,
updatedContent: string,
isSequence: boolean,
isInput: boolean
): IStep[] => {
let hasChanges = false;
const nextMessages = messages.map((msg) => {
if (isEqual(msg.id, messageId)) {
hasChanges = true;
const newMsg = { ...msg };
if ('content' in newMsg && newMsg.content !== undefined) {
if (isSequence) {
newMsg.content = updatedContent;
} else {
newMsg.content += updatedContent;
}
} else if (isInput) {
if ('input' in newMsg && newMsg.input !== undefined) {
if (isSequence) {
newMsg.input = updatedContent;
} else {
newMsg.input += updatedContent;
}
}
} else {
if ('output' in newMsg && newMsg.output !== undefined) {
if (isSequence) {
newMsg.output = updatedContent;
} else {
newMsg.output += updatedContent;
}
}
}
return newMsg;
} else if (msg.steps) {
const updatedSteps = updateMessageContentById(
msg.steps,
messageId,
updatedContent,
isSequence,
isInput
);
if (updatedSteps !== msg.steps) {
hasChanges = true;
return { ...msg, steps: updatedSteps };
}
}
return msg;
});
return hasChanges ? nextMessages : messages;
};
export {
addMessageToParent,
addMessage,
deleteMessageById,
hasMessageById,
isLastMessage,
nestMessages,
updateMessageById,
updateMessageContentById
};
@@ -0,0 +1,203 @@
import {
noteFrequencies,
noteFrequencyLabels,
voiceFrequencies,
voiceFrequencyLabels
} from './constants.js';
/**
* Output of AudioAnalysis for the frequency domain of the audio
* @typedef {Object} AudioAnalysisOutputType
* @property {Float32Array} values Amplitude of this frequency between {0, 1} inclusive
* @property {number[]} frequencies Raw frequency bucket values
* @property {string[]} labels Labels for the frequency bucket values
*/
/**
* Analyzes audio for visual output
* @class
*/
export class AudioAnalysis {
/**
* Retrieves frequency domain data from an AnalyserNode adjusted to a decibel range
* returns human-readable formatting and labels
* @param {AnalyserNode} analyser
* @param {number} sampleRate
* @param {Float32Array} [fftResult]
* @param {"frequency"|"music"|"voice"} [analysisType]
* @param {number} [minDecibels] default -100
* @param {number} [maxDecibels] default -30
* @returns {AudioAnalysisOutputType}
*/
static getFrequencies(
analyser,
sampleRate,
fftResult,
analysisType = 'frequency',
minDecibels = -100,
maxDecibels = -30
) {
if (!fftResult) {
fftResult = new Float32Array(analyser.frequencyBinCount);
analyser.getFloatFrequencyData(fftResult);
}
const nyquistFrequency = sampleRate / 2;
const frequencyStep = (1 / fftResult.length) * nyquistFrequency;
let outputValues;
let frequencies;
let labels;
if (analysisType === 'music' || analysisType === 'voice') {
const useFrequencies =
analysisType === 'voice' ? voiceFrequencies : noteFrequencies;
const aggregateOutput = Array(useFrequencies.length).fill(minDecibels);
for (let i = 0; i < fftResult.length; i++) {
const frequency = i * frequencyStep;
const amplitude = fftResult[i];
for (let n = useFrequencies.length - 1; n >= 0; n--) {
if (frequency > useFrequencies[n]) {
aggregateOutput[n] = Math.max(aggregateOutput[n], amplitude);
break;
}
}
}
outputValues = aggregateOutput;
frequencies =
analysisType === 'voice' ? voiceFrequencies : noteFrequencies;
labels =
analysisType === 'voice' ? voiceFrequencyLabels : noteFrequencyLabels;
} else {
outputValues = Array.from(fftResult);
frequencies = outputValues.map((_, i) => frequencyStep * i);
labels = frequencies.map((f) => `${f.toFixed(2)} Hz`);
}
// We normalize to {0, 1}
const normalizedOutput = outputValues.map((v) => {
return Math.max(
0,
Math.min((v - minDecibels) / (maxDecibels - minDecibels), 1)
);
});
const values = new Float32Array(normalizedOutput);
return {
values,
frequencies,
labels
};
}
/**
* Creates a new AudioAnalysis instance for an HTMLAudioElement
* @param {HTMLAudioElement} audioElement
* @param {AudioBuffer|null} [audioBuffer] If provided, will cache all frequency domain data from the buffer
* @returns {AudioAnalysis}
*/
constructor(audioElement, audioBuffer = null) {
this.fftResults = [];
if (audioBuffer) {
/**
* Modified from
* https://stackoverflow.com/questions/75063715/using-the-web-audio-api-to-analyze-a-song-without-playing
*
* We do this to populate FFT values for the audio if provided an `audioBuffer`
* The reason to do this is that Safari fails when using `createMediaElementSource`
* This has a non-zero RAM cost so we only opt-in to run it on Safari, Chrome is better
*/
const { length, sampleRate } = audioBuffer;
const offlineAudioContext = new OfflineAudioContext({
length,
sampleRate
});
const source = offlineAudioContext.createBufferSource();
source.buffer = audioBuffer;
const analyser = offlineAudioContext.createAnalyser();
analyser.fftSize = 8192;
analyser.smoothingTimeConstant = 0.1;
source.connect(analyser);
// limit is :: 128 / sampleRate;
// but we just want 60fps - cuts ~1s from 6MB to 1MB of RAM
const renderQuantumInSeconds = 1 / 60;
const durationInSeconds = length / sampleRate;
const analyze = (index) => {
const suspendTime = renderQuantumInSeconds * index;
if (suspendTime < durationInSeconds) {
offlineAudioContext.suspend(suspendTime).then(() => {
const fftResult = new Float32Array(analyser.frequencyBinCount);
analyser.getFloatFrequencyData(fftResult);
this.fftResults.push(fftResult);
analyze(index + 1);
});
}
if (index === 1) {
offlineAudioContext.startRendering();
} else {
offlineAudioContext.resume();
}
};
source.start(0);
analyze(1);
this.audio = audioElement;
this.context = offlineAudioContext;
this.analyser = analyser;
this.sampleRate = sampleRate;
this.audioBuffer = audioBuffer;
} else {
const audioContext = new AudioContext();
const track = audioContext.createMediaElementSource(audioElement);
const analyser = audioContext.createAnalyser();
analyser.fftSize = 8192;
analyser.smoothingTimeConstant = 0.1;
track.connect(analyser);
analyser.connect(audioContext.destination);
this.audio = audioElement;
this.context = audioContext;
this.analyser = analyser;
this.sampleRate = this.context.sampleRate;
this.audioBuffer = null;
}
}
/**
* Gets the current frequency domain data from the playing audio track
* @param {"frequency"|"music"|"voice"} [analysisType]
* @param {number} [minDecibels] default -100
* @param {number} [maxDecibels] default -30
* @returns {AudioAnalysisOutputType}
*/
getFrequencies(
analysisType = 'frequency',
minDecibels = -100,
maxDecibels = -30
) {
let fftResult = null;
if (this.audioBuffer && this.fftResults.length) {
const pct = this.audio.currentTime / this.audio.duration;
const index = Math.min(
(pct * this.fftResults.length) | 0,
this.fftResults.length - 1
);
fftResult = this.fftResults[index];
}
return AudioAnalysis.getFrequencies(
this.analyser,
this.sampleRate,
fftResult,
analysisType,
minDecibels,
maxDecibels
);
}
/**
* Resume the internal AudioContext if it was suspended due to the lack of
* user interaction when the AudioAnalysis was instantiated.
* @returns {Promise<true>}
*/
async resumeIfSuspended() {
if (this.context.state === 'suspended') {
await this.context.resume();
}
return true;
}
}
globalThis.AudioAnalysis = AudioAnalysis;
+60
View File
@@ -0,0 +1,60 @@
/**
* Constants for help with visualization
* Helps map frequency ranges from Fast Fourier Transform
* to human-interpretable ranges, notably music ranges and
* human vocal ranges.
*/
// Eighth octave frequencies
const octave8Frequencies = [
4186.01, 4434.92, 4698.63, 4978.03, 5274.04, 5587.65, 5919.91, 6271.93,
6644.88, 7040.0, 7458.62, 7902.13
];
// Labels for each of the above frequencies
const octave8FrequencyLabels = [
'C',
'C#',
'D',
'D#',
'E',
'F',
'F#',
'G',
'G#',
'A',
'A#',
'B'
];
/**
* All note frequencies from 1st to 8th octave
* in format "A#8" (A#, 8th octave)
*/
export const noteFrequencies = [];
export const noteFrequencyLabels = [];
for (let i = 1; i <= 8; i++) {
for (let f = 0; f < octave8Frequencies.length; f++) {
const freq = octave8Frequencies[f];
noteFrequencies.push(freq / Math.pow(2, 8 - i));
noteFrequencyLabels.push(octave8FrequencyLabels[f] + i);
}
}
/**
* Subset of the note frequencies between 32 and 2000 Hz
* 6 octave range: C1 to B6
*/
const voiceFrequencyRange = [32.0, 2000.0];
export const voiceFrequencies = noteFrequencies.filter((_, i) => {
return (
noteFrequencies[i] > voiceFrequencyRange[0] &&
noteFrequencies[i] < voiceFrequencyRange[1]
);
});
export const voiceFrequencyLabels = noteFrequencyLabels.filter((_, i) => {
return (
noteFrequencies[i] > voiceFrequencyRange[0] &&
noteFrequencies[i] < voiceFrequencyRange[1]
);
});
+7
View File
@@ -0,0 +1,7 @@
// Courtesy of https://github.com/openai/openai-realtime-console
import { AudioAnalysis } from './analysis/audio_analysis.js';
import { WavPacker } from './wav_packer.js';
import { WavRecorder } from './wav_recorder.js';
import { WavStreamPlayer } from './wav_stream_player.js';
export { AudioAnalysis, WavPacker, WavStreamPlayer, WavRecorder };
+113
View File
@@ -0,0 +1,113 @@
/**
* Raw wav audio file contents
* @typedef {Object} WavPackerAudioType
* @property {Blob} blob
* @property {string} url
* @property {number} channelCount
* @property {number} sampleRate
* @property {number} duration
*/
/**
* Utility class for assembling PCM16 "audio/wav" data
* @class
*/
export class WavPacker {
/**
* Converts Float32Array of amplitude data to ArrayBuffer in Int16Array format
* @param {Float32Array} float32Array
* @returns {ArrayBuffer}
*/
static floatTo16BitPCM(float32Array) {
const buffer = new ArrayBuffer(float32Array.length * 2);
const view = new DataView(buffer);
let offset = 0;
for (let i = 0; i < float32Array.length; i++, offset += 2) {
let s = Math.max(-1, Math.min(1, float32Array[i]));
view.setInt16(offset, s < 0 ? s * 0x8000 : s * 0x7fff, true);
}
return buffer;
}
/**
* Concatenates two ArrayBuffers
* @param {ArrayBuffer} leftBuffer
* @param {ArrayBuffer} rightBuffer
* @returns {ArrayBuffer}
*/
static mergeBuffers(leftBuffer, rightBuffer) {
const tmpArray = new Uint8Array(
leftBuffer.byteLength + rightBuffer.byteLength
);
tmpArray.set(new Uint8Array(leftBuffer), 0);
tmpArray.set(new Uint8Array(rightBuffer), leftBuffer.byteLength);
return tmpArray.buffer;
}
/**
* Packs data into an Int16 format
* @private
* @param {number} size 0 = 1x Int16, 1 = 2x Int16
* @param {number} arg value to pack
* @returns
*/
_packData(size, arg) {
return [
new Uint8Array([arg, arg >> 8]),
new Uint8Array([arg, arg >> 8, arg >> 16, arg >> 24])
][size];
}
/**
* Packs audio into "audio/wav" Blob
* @param {number} sampleRate
* @param {{bitsPerSample: number, channels: Array<Float32Array>, data: Int16Array}} audio
* @returns {WavPackerAudioType}
*/
pack(sampleRate, audio) {
if (!audio?.bitsPerSample) {
throw new Error(`Missing "bitsPerSample"`);
} else if (!audio?.channels) {
throw new Error(`Missing "channels"`);
} else if (!audio?.data) {
throw new Error(`Missing "data"`);
}
const { bitsPerSample, channels, data } = audio;
const output = [
// Header
'RIFF',
this._packData(
1,
4 + (8 + 24) /* chunk 1 length */ + (8 + 8) /* chunk 2 length */
), // Length
'WAVE',
// chunk 1
'fmt ', // Sub-chunk identifier
this._packData(1, 16), // Chunk length
this._packData(0, 1), // Audio format (1 is linear quantization)
this._packData(0, channels.length),
this._packData(1, sampleRate),
this._packData(1, (sampleRate * channels.length * bitsPerSample) / 8), // Byte rate
this._packData(0, (channels.length * bitsPerSample) / 8),
this._packData(0, bitsPerSample),
// chunk 2
'data', // Sub-chunk identifier
this._packData(
1,
(channels[0].length * channels.length * bitsPerSample) / 8
), // Chunk length
data
];
const blob = new Blob(output, { type: 'audio/mpeg' });
const url = URL.createObjectURL(blob);
return {
blob,
url,
channelCount: channels.length,
sampleRate,
duration: data.byteLength / (channels.length * sampleRate * 2)
};
}
}
globalThis.WavPacker = WavPacker;
+549
View File
@@ -0,0 +1,549 @@
import { AudioAnalysis } from './analysis/audio_analysis.js';
import { WavPacker } from './wav_packer.js';
import { AudioProcessorSrc } from './worklets/audio_processor.js';
/**
* Decodes audio into a wav file
* @typedef {Object} DecodedAudioType
* @property {Blob} blob
* @property {string} url
* @property {Float32Array} values
* @property {AudioBuffer} audioBuffer
*/
/**
* Records live stream of user audio as PCM16 "audio/wav" data
* @class
*/
export class WavRecorder {
/**
* Create a new WavRecorder instance
* @param {{sampleRate?: number, outputToSpeakers?: boolean, debug?: boolean}} [options]
* @returns {WavRecorder}
*/
constructor({
sampleRate = 24000,
outputToSpeakers = false,
debug = false
} = {}) {
// Script source
this.scriptSrc = AudioProcessorSrc;
// Config
this.sampleRate = sampleRate;
this.outputToSpeakers = outputToSpeakers;
this.debug = !!debug;
this._deviceChangeCallback = null;
this._devices = [];
// State variables
this.stream = null;
this.processor = null;
this.source = null;
this.node = null;
this.recording = false;
// Event handling with AudioWorklet
this._lastEventId = 0;
this.eventReceipts = {};
this.eventTimeout = 5000;
// Process chunks of audio
this._chunkProcessor = () => {};
this._chunkProcessorSize = void 0;
this._chunkProcessorBuffer = {
raw: new ArrayBuffer(0),
mono: new ArrayBuffer(0)
};
}
/**
* Decodes audio data from multiple formats to a Blob, url, Float32Array and AudioBuffer
* @param {Blob|Float32Array|Int16Array|ArrayBuffer|number[]} audioData
* @param {number} sampleRate
* @param {number} fromSampleRate
* @returns {Promise<DecodedAudioType>}
*/
static async decode(audioData, sampleRate = 24000, fromSampleRate = -1) {
const context = new AudioContext({ sampleRate });
let arrayBuffer;
let blob;
if (audioData instanceof Blob) {
if (fromSampleRate !== -1) {
throw new Error(
`Can not specify "fromSampleRate" when reading from Blob`
);
}
blob = audioData;
arrayBuffer = await blob.arrayBuffer();
} else if (audioData instanceof ArrayBuffer) {
if (fromSampleRate !== -1) {
throw new Error(
`Can not specify "fromSampleRate" when reading from ArrayBuffer`
);
}
arrayBuffer = audioData;
blob = new Blob([arrayBuffer], { type: 'audio/wav' });
} else {
let float32Array;
let data;
if (audioData instanceof Int16Array) {
data = audioData;
float32Array = new Float32Array(audioData.length);
for (let i = 0; i < audioData.length; i++) {
float32Array[i] = audioData[i] / 0x8000;
}
} else if (audioData instanceof Float32Array) {
float32Array = audioData;
} else if (audioData instanceof Array) {
float32Array = new Float32Array(audioData);
} else {
throw new Error(
`"audioData" must be one of: Blob, Float32Arrray, Int16Array, ArrayBuffer, Array<number>`
);
}
if (fromSampleRate === -1) {
throw new Error(
`Must specify "fromSampleRate" when reading from Float32Array, In16Array or Array`
);
} else if (fromSampleRate < 3000) {
throw new Error(`Minimum "fromSampleRate" is 3000 (3kHz)`);
}
if (!data) {
data = WavPacker.floatTo16BitPCM(float32Array);
}
const audio = {
bitsPerSample: 16,
channels: [float32Array],
data
};
const packer = new WavPacker();
const result = packer.pack(fromSampleRate, audio);
blob = result.blob;
arrayBuffer = await blob.arrayBuffer();
}
const audioBuffer = await context.decodeAudioData(arrayBuffer);
const values = audioBuffer.getChannelData(0);
const url = URL.createObjectURL(blob);
return {
blob,
url,
values,
audioBuffer
};
}
/**
* Logs data in debug mode
* @param {...any} arguments
* @returns {true}
*/
log() {
if (this.debug) {
this.log(...arguments);
}
return true;
}
/**
* Retrieves the current sampleRate for the recorder
* @returns {number}
*/
getSampleRate() {
return this.sampleRate;
}
/**
* Retrieves the current status of the recording
* @returns {"ended"|"paused"|"recording"}
*/
getStatus() {
if (!this.processor) {
return 'ended';
} else if (!this.recording) {
return 'paused';
} else {
return 'recording';
}
}
/**
* Sends an event to the AudioWorklet
* @private
* @param {string} name
* @param {{[key: string]: any}} data
* @param {AudioWorkletNode} [_processor]
* @returns {Promise<{[key: string]: any}>}
*/
async _event(name, data = {}, _processor = null) {
_processor = _processor || this.processor;
if (!_processor) {
throw new Error('Can not send events without recording first');
}
const message = {
event: name,
id: this._lastEventId++,
data
};
_processor.port.postMessage(message);
const t0 = new Date().valueOf();
while (!this.eventReceipts[message.id]) {
if (new Date().valueOf() - t0 > this.eventTimeout) {
throw new Error(`Timeout waiting for "${name}" event`);
}
await new Promise((res) => setTimeout(() => res(true), 1));
}
const payload = this.eventReceipts[message.id];
delete this.eventReceipts[message.id];
return payload;
}
/**
* Sets device change callback, remove if callback provided is `null`
* @param {(Array<MediaDeviceInfo & {default: boolean}>): void|null} callback
* @returns {true}
*/
listenForDeviceChange(callback) {
if (callback === null && this._deviceChangeCallback) {
navigator.mediaDevices.removeEventListener(
'devicechange',
this._deviceChangeCallback
);
this._deviceChangeCallback = null;
} else if (callback !== null) {
// Basically a debounce; we only want this called once when devices change
// And we only want the most recent callback() to be executed
// if a few are operating at the same time
let lastId = 0;
let lastDevices = [];
const serializeDevices = (devices) =>
devices
.map((d) => d.deviceId)
.sort()
.join(',');
const cb = async () => {
let id = ++lastId;
const devices = await this.listDevices();
if (id === lastId) {
if (serializeDevices(lastDevices) !== serializeDevices(devices)) {
lastDevices = devices;
callback(devices.slice());
}
}
};
navigator.mediaDevices.addEventListener('devicechange', cb);
cb();
this._deviceChangeCallback = cb;
}
return true;
}
/**
* Manually request permission to use the microphone
* @returns {Promise<true>}
*/
async requestPermission() {
const permissionStatus = await navigator.permissions.query({
name: 'microphone'
});
if (permissionStatus.state === 'denied') {
window.alert('You must grant microphone access to use this feature.');
} else if (permissionStatus.state === 'prompt') {
try {
const stream = await navigator.mediaDevices.getUserMedia({
audio: true
});
const tracks = stream.getTracks();
tracks.forEach((track) => track.stop());
} catch (_e) {
window.alert('You must grant microphone access to use this feature.');
}
}
return true;
}
/**
* List all eligible devices for recording, will request permission to use microphone
* @returns {Promise<Array<MediaDeviceInfo & {default: boolean}>>}
*/
async listDevices() {
if (
!navigator.mediaDevices ||
!('enumerateDevices' in navigator.mediaDevices)
) {
throw new Error('Could not request user devices');
}
await this.requestPermission();
const devices = await navigator.mediaDevices.enumerateDevices();
const audioDevices = devices.filter(
(device) => device.kind === 'audioinput'
);
const defaultDeviceIndex = audioDevices.findIndex(
(device) => device.deviceId === 'default'
);
const deviceList = [];
if (defaultDeviceIndex !== -1) {
let defaultDevice = audioDevices.splice(defaultDeviceIndex, 1)[0];
let existingIndex = audioDevices.findIndex(
(device) => device.groupId === defaultDevice.groupId
);
if (existingIndex !== -1) {
defaultDevice = audioDevices.splice(existingIndex, 1)[0];
}
defaultDevice.default = true;
deviceList.push(defaultDevice);
}
return deviceList.concat(audioDevices);
}
/**
* Begins a recording session and requests microphone permissions if not already granted
* Microphone recording indicator will appear on browser tab but status will be "paused"
* @param {string} [deviceId] if no device provided, default device will be used
* @returns {Promise<true>}
*/
async begin(deviceId) {
if (this.processor) {
throw new Error(
`Already connected: please call .end() to start a new session`
);
}
if (
!navigator.mediaDevices ||
!('getUserMedia' in navigator.mediaDevices)
) {
throw new Error('Could not request user media');
}
try {
const config = { audio: true };
if (deviceId) {
config.audio = { deviceId: { exact: deviceId } };
}
this.stream = await navigator.mediaDevices.getUserMedia(config);
} catch (err) {
throw new Error('Could not start media stream', { cause: err });
}
const context = new AudioContext({ sampleRate: this.sampleRate });
const source = context.createMediaStreamSource(this.stream);
// Load and execute the module script.
try {
await context.audioWorklet.addModule(this.scriptSrc);
} catch (e) {
console.error(e);
throw new Error(`Could not add audioWorklet module: ${this.scriptSrc}`, {
cause: e
});
}
const processor = new AudioWorkletNode(context, 'audio_processor');
processor.port.onmessage = (e) => {
const { event, id, data } = e.data;
if (event === 'receipt') {
this.eventReceipts[id] = data;
} else if (event === 'chunk') {
if (this._chunkProcessorSize) {
const buffer = this._chunkProcessorBuffer;
this._chunkProcessorBuffer = {
raw: WavPacker.mergeBuffers(buffer.raw, data.raw),
mono: WavPacker.mergeBuffers(buffer.mono, data.mono)
};
if (
this._chunkProcessorBuffer.mono.byteLength >=
this._chunkProcessorSize
) {
this._chunkProcessor(this._chunkProcessorBuffer);
this._chunkProcessorBuffer = {
raw: new ArrayBuffer(0),
mono: new ArrayBuffer(0)
};
}
} else {
this._chunkProcessor(data);
}
}
};
const node = source.connect(processor);
const analyser = context.createAnalyser();
analyser.fftSize = 8192;
analyser.smoothingTimeConstant = 0.1;
node.connect(analyser);
if (this.outputToSpeakers) {
console.warn(
'Warning: Output to speakers may affect sound quality,\n' +
'especially due to system audio feedback preventative measures.\n' +
'use only for debugging'
);
analyser.connect(context.destination);
}
this.source = source;
this.node = node;
this.analyser = analyser;
this.processor = processor;
return true;
}
/**
* Gets the current frequency domain data from the recording track
* @param {"frequency"|"music"|"voice"} [analysisType]
* @param {number} [minDecibels] default -100
* @param {number} [maxDecibels] default -30
* @returns {import('./analysis/audio_analysis.js').AudioAnalysisOutputType}
*/
getFrequencies(
analysisType = 'frequency',
minDecibels = -100,
maxDecibels = -30
) {
if (!this.processor) {
throw new Error('Session ended: please call .begin() first');
}
return AudioAnalysis.getFrequencies(
this.analyser,
this.sampleRate,
null,
analysisType,
minDecibels,
maxDecibels
);
}
/**
* Pauses the recording
* Keeps microphone stream open but halts storage of audio
* @returns {Promise<true>}
*/
async pause() {
if (!this.processor) {
throw new Error('Session ended: please call .begin() first');
} else if (!this.recording) {
throw new Error('Already paused: please call .record() first');
}
if (this._chunkProcessorBuffer.raw.byteLength) {
this._chunkProcessor(this._chunkProcessorBuffer);
}
this.log('Pausing ...');
await this._event('stop');
this.recording = false;
return true;
}
/**
* Start recording stream and storing to memory from the connected audio source
* @param {(data: { mono: Int16Array; raw: Int16Array }) => any} [chunkProcessor]
* @param {number} [chunkSize] chunkProcessor will not be triggered until this size threshold met in mono audio
* @returns {Promise<true>}
*/
async record(chunkProcessor = () => {}, chunkSize = 8192) {
if (!this.processor) {
throw new Error('Session ended: please call .begin() first');
} else if (this.recording) {
throw new Error('Already recording: please call .pause() first');
} else if (typeof chunkProcessor !== 'function') {
throw new Error(`chunkProcessor must be a function`);
}
this._chunkProcessor = chunkProcessor;
this._chunkProcessorSize = chunkSize;
this._chunkProcessorBuffer = {
raw: new ArrayBuffer(0),
mono: new ArrayBuffer(0)
};
this.log('Recording ...');
await this._event('start');
this.recording = true;
return true;
}
/**
* Clears the audio buffer, empties stored recording
* @returns {Promise<true>}
*/
async clear() {
if (!this.processor) {
throw new Error('Session ended: please call .begin() first');
}
await this._event('clear');
return true;
}
/**
* Reads the current audio stream data
* @returns {Promise<{meanValues: Float32Array, channels: Array<Float32Array>}>}
*/
async read() {
if (!this.processor) {
throw new Error('Session ended: please call .begin() first');
}
this.log('Reading ...');
const result = await this._event('read');
return result;
}
/**
* Saves the current audio stream to a file
* @param {boolean} [force] Force saving while still recording
* @returns {Promise<import('./wav_packer.js').WavPackerAudioType>}
*/
async save(force = false) {
if (!this.processor) {
throw new Error('Session ended: please call .begin() first');
}
if (!force && this.recording) {
throw new Error(
'Currently recording: please call .pause() first, or call .save(true) to force'
);
}
this.log('Exporting ...');
const exportData = await this._event('export');
const packer = new WavPacker();
const result = packer.pack(this.sampleRate, exportData.audio);
return result;
}
/**
* Ends the current recording session and saves the result
* @returns {Promise<import('./wav_packer.js').WavPackerAudioType>}
*/
async end() {
if (!this.processor) {
throw new Error('Session ended: please call .begin() first');
}
const _processor = this.processor;
this.log('Stopping ...');
await this._event('stop');
this.recording = false;
const tracks = this.stream.getTracks();
tracks.forEach((track) => track.stop());
this.log('Exporting ...');
const exportData = await this._event('export', {}, _processor);
this.processor.disconnect();
this.source.disconnect();
this.node.disconnect();
this.analyser.disconnect();
this.stream = null;
this.processor = null;
this.source = null;
this.node = null;
const packer = new WavPacker();
const result = packer.pack(this.sampleRate, exportData.audio);
return result;
}
/**
* Performs a full cleanup of WavRecorder instance
* Stops actively listening via microphone and removes existing listeners
* @returns {Promise<true>}
*/
async quit() {
this.listenForDeviceChange(null);
if (this.processor) {
await this.end();
}
return true;
}
}
globalThis.WavRecorder = WavRecorder;
@@ -0,0 +1,132 @@
const dataMap = new WeakMap();
/**
* Normalizes a Float32Array to Array(m): We use this to draw amplitudes on a graph
* If we're rendering the same audio data, then we'll often be using
* the same (data, m, downsamplePeaks) triplets so we give option to memoize
*/
const normalizeArray = (
data: Float32Array,
m: number,
downsamplePeaks: boolean = false,
memoize: boolean = false
) => {
let cache, mKey, dKey;
if (memoize) {
mKey = m.toString();
dKey = downsamplePeaks.toString();
cache = dataMap.has(data) ? dataMap.get(data) : {};
dataMap.set(data, cache);
cache[mKey] = cache[mKey] || {};
if (cache[mKey][dKey]) {
return cache[mKey][dKey];
}
}
const n = data.length;
const result = new Array(m);
if (m <= n) {
// Downsampling
result.fill(0);
const count = new Array(m).fill(0);
for (let i = 0; i < n; i++) {
const index = Math.floor(i * (m / n));
if (downsamplePeaks) {
// take highest result in the set
result[index] = Math.max(result[index], Math.abs(data[i]));
} else {
result[index] += Math.abs(data[i]);
}
count[index]++;
}
if (!downsamplePeaks) {
for (let i = 0; i < result.length; i++) {
result[i] = result[i] / count[i];
}
}
} else {
for (let i = 0; i < m; i++) {
const index = (i * (n - 1)) / (m - 1);
const low = Math.floor(index);
const high = Math.ceil(index);
const t = index - low;
if (high >= n) {
result[i] = data[n - 1];
} else {
result[i] = data[low] * (1 - t) + data[high] * t;
}
}
}
if (memoize) {
cache[mKey as string][dKey as string] = result;
}
return result;
};
export const WavRenderer = {
/**
* Renders a point-in-time snapshot of an audio sample, usually frequency values
* @param ctx
* @param data
* @param color
* @param cssWidth
* @param cssHeight
* @param pointCount number of bars to render
* @param barWidth width of bars in px
* @param barSpacing spacing between bars in px
* @param center vertically center the bars
*/
drawBars: (
ctx: CanvasRenderingContext2D,
data: Float32Array,
cssWidth: number,
cssHeight: number,
color: string,
pointCount: number = 0,
barWidth: number = 0,
barSpacing: number = 0,
center: boolean = false
) => {
pointCount = Math.floor(
Math.min(
pointCount,
(cssWidth - barSpacing) / (Math.max(barWidth, 1) + barSpacing)
)
);
if (!pointCount) {
pointCount = Math.floor(
(cssWidth - barSpacing) / (Math.max(barWidth, 1) + barSpacing)
);
}
if (!barWidth) {
barWidth = (cssWidth - barSpacing) / pointCount - barSpacing;
}
const points = normalizeArray(data, pointCount, true);
for (let i = 0; i < pointCount; i++) {
const amplitude = Math.abs(points[i]);
const height = Math.max(1, amplitude * cssHeight);
const x = barSpacing + i * (barWidth + barSpacing);
const y = center ? (cssHeight - height) / 2 : cssHeight - height;
const radius = Math.min(barWidth / 2, height / 2); // Calculate the radius for rounded corners
ctx.fillStyle = color;
ctx.beginPath();
ctx.moveTo(x + radius, y);
ctx.lineTo(x + barWidth - radius, y);
ctx.arcTo(x + barWidth, y, x + barWidth, y + radius, radius);
ctx.lineTo(x + barWidth, y + height - radius);
ctx.arcTo(
x + barWidth,
y + height,
x + barWidth - radius,
y + height,
radius
);
ctx.lineTo(x + radius, y + height);
ctx.arcTo(x, y + height, x, y + height - radius, radius);
ctx.lineTo(x, y + radius);
ctx.arcTo(x, y, x + radius, y, radius);
ctx.closePath();
ctx.fill();
}
}
};
+164
View File
@@ -0,0 +1,164 @@
import { AudioAnalysis } from './analysis/audio_analysis.js';
import { StreamProcessorSrc } from './worklets/stream_processor.js';
/**
* Plays audio streams received in raw PCM16 chunks from the browser
* @class
*/
export class WavStreamPlayer {
/**
* Creates a new WavStreamPlayer instance
* @param {{sampleRate?: number}} options
* @returns {WavStreamPlayer}
*/
constructor({ sampleRate = 24000, onStop } = {}) {
this.scriptSrc = StreamProcessorSrc;
this.onStop = onStop;
this.sampleRate = sampleRate;
this.context = null;
this.stream = null;
this.analyser = null;
this.trackSampleOffsets = {};
this.interruptedTrackIds = {};
}
/**
* Connects the audio context and enables output to speakers
* @returns {Promise<true>}
*/
async connect() {
this.context = new AudioContext({ sampleRate: this.sampleRate });
if (this.context.state === 'suspended') {
await this.context.resume();
}
try {
await this.context.audioWorklet.addModule(this.scriptSrc);
} catch (e) {
console.error(e);
throw new Error(`Could not add audioWorklet module: ${this.scriptSrc}`, {
cause: e
});
}
const analyser = this.context.createAnalyser();
analyser.fftSize = 8192;
analyser.smoothingTimeConstant = 0.1;
this.analyser = analyser;
return true;
}
/**
* Gets the current frequency domain data from the playing track
* @param {"frequency"|"music"|"voice"} [analysisType]
* @param {number} [minDecibels] default -100
* @param {number} [maxDecibels] default -30
* @returns {import('./analysis/audio_analysis.js').AudioAnalysisOutputType}
*/
getFrequencies(
analysisType = 'frequency',
minDecibels = -100,
maxDecibels = -30
) {
if (!this.analyser) {
throw new Error('Not connected, please call .connect() first');
}
return AudioAnalysis.getFrequencies(
this.analyser,
this.sampleRate,
null,
analysisType,
minDecibels,
maxDecibels
);
}
/**
* Starts audio streaming
* @private
* @returns {Promise<true>}
*/
_start() {
const streamNode = new AudioWorkletNode(this.context, 'stream_processor');
streamNode.connect(this.context.destination);
streamNode.port.onmessage = (e) => {
const { event } = e.data;
if (event === 'stop') {
this.onStop?.();
streamNode.disconnect();
this.stream = null;
} else if (event === 'offset') {
const { requestId, trackId, offset } = e.data;
const currentTime = offset / this.sampleRate;
this.trackSampleOffsets[requestId] = { trackId, offset, currentTime };
}
};
this.analyser.disconnect();
streamNode.connect(this.analyser);
this.stream = streamNode;
return true;
}
/**
* Adds 16BitPCM data to the currently playing audio stream
* You can add chunks beyond the current play point and they will be queued for play
* @param {ArrayBuffer|Int16Array} arrayBuffer
* @param {string} [trackId]
* @returns {Int16Array}
*/
add16BitPCM(arrayBuffer, trackId = 'default') {
if (typeof trackId !== 'string') {
throw new Error(`trackId must be a string`);
} else if (this.interruptedTrackIds[trackId]) {
return;
}
if (!this.stream) {
this._start();
}
let buffer;
if (arrayBuffer instanceof Int16Array) {
buffer = arrayBuffer;
} else if (arrayBuffer instanceof ArrayBuffer) {
buffer = new Int16Array(arrayBuffer);
} else {
throw new Error(`argument must be Int16Array or ArrayBuffer`);
}
this.stream.port.postMessage({ event: 'write', buffer, trackId });
return buffer;
}
/**
* Gets the offset (sample count) of the currently playing stream
* @param {boolean} [interrupt]
* @returns {{trackId: string|null, offset: number, currentTime: number}}
*/
async getTrackSampleOffset(interrupt = false) {
if (!this.stream) {
return null;
}
const requestId = crypto.randomUUID();
this.stream.port.postMessage({
event: interrupt ? 'interrupt' : 'offset',
requestId
});
let trackSampleOffset;
while (!trackSampleOffset) {
trackSampleOffset = this.trackSampleOffsets[requestId];
await new Promise((r) => setTimeout(() => r(), 1));
}
const { trackId } = trackSampleOffset;
if (interrupt && trackId) {
this.interruptedTrackIds[trackId] = true;
}
return trackSampleOffset;
}
/**
* Strips the current stream and returns the sample offset of the audio
* @param {boolean} [interrupt]
* @returns {{trackId: string|null, offset: number, currentTime: number}}
*/
async interrupt() {
return this.getTrackSampleOffset(true);
}
}
globalThis.WavStreamPlayer = WavStreamPlayer;
@@ -0,0 +1,214 @@
const AudioProcessorWorklet = `
class AudioProcessor extends AudioWorkletProcessor {
constructor() {
super();
this.port.onmessage = this.receive.bind(this);
this.initialize();
}
initialize() {
this.foundAudio = false;
this.recording = false;
this.chunks = [];
}
/**
* Concatenates sampled chunks into channels
* Format is chunk[Left[], Right[]]
*/
readChannelData(chunks, channel = -1, maxChannels = 9) {
let channelLimit;
if (channel !== -1) {
if (chunks[0] && chunks[0].length - 1 < channel) {
throw new Error(
\`Channel \${channel} out of range: max \${chunks[0].length}\`
);
}
channelLimit = channel + 1;
} else {
channel = 0;
channelLimit = Math.min(chunks[0] ? chunks[0].length : 1, maxChannels);
}
const channels = [];
for (let n = channel; n < channelLimit; n++) {
const length = chunks.reduce((sum, chunk) => {
return sum + chunk[n].length;
}, 0);
const buffers = chunks.map((chunk) => chunk[n]);
const result = new Float32Array(length);
let offset = 0;
for (let i = 0; i < buffers.length; i++) {
result.set(buffers[i], offset);
offset += buffers[i].length;
}
channels[n] = result;
}
return channels;
}
/**
* Combines parallel audio data into correct format,
* channels[Left[], Right[]] to float32Array[LRLRLRLR...]
*/
formatAudioData(channels) {
if (channels.length === 1) {
// Simple case is only one channel
const float32Array = channels[0].slice();
const meanValues = channels[0].slice();
return { float32Array, meanValues };
} else {
const float32Array = new Float32Array(
channels[0].length * channels.length
);
const meanValues = new Float32Array(channels[0].length);
for (let i = 0; i < channels[0].length; i++) {
const offset = i * channels.length;
let meanValue = 0;
for (let n = 0; n < channels.length; n++) {
float32Array[offset + n] = channels[n][i];
meanValue += channels[n][i];
}
meanValues[i] = meanValue / channels.length;
}
return { float32Array, meanValues };
}
}
/**
* Converts 32-bit float data to 16-bit integers
*/
floatTo16BitPCM(float32Array) {
const buffer = new ArrayBuffer(float32Array.length * 2);
const view = new DataView(buffer);
let offset = 0;
for (let i = 0; i < float32Array.length; i++, offset += 2) {
let s = Math.max(-1, Math.min(1, float32Array[i]));
view.setInt16(offset, s < 0 ? s * 0x8000 : s * 0x7fff, true);
}
return buffer;
}
/**
* Retrieves the most recent amplitude values from the audio stream
* @param {number} channel
*/
getValues(channel = -1) {
const channels = this.readChannelData(this.chunks, channel);
const { meanValues } = this.formatAudioData(channels);
return { meanValues, channels };
}
/**
* Exports chunks as an audio/wav file
*/
export() {
const channels = this.readChannelData(this.chunks);
const { float32Array, meanValues } = this.formatAudioData(channels);
const audioData = this.floatTo16BitPCM(float32Array);
return {
meanValues: meanValues,
audio: {
bitsPerSample: 16,
channels: channels,
data: audioData,
},
};
}
receive(e) {
const { event, id } = e.data;
let receiptData = {};
switch (event) {
case 'start':
this.recording = true;
break;
case 'stop':
this.recording = false;
break;
case 'clear':
this.initialize();
break;
case 'export':
receiptData = this.export();
break;
case 'read':
receiptData = this.getValues();
break;
default:
break;
}
// Always send back receipt
this.port.postMessage({ event: 'receipt', id, data: receiptData });
}
sendChunk(chunk) {
const channels = this.readChannelData([chunk]);
const { float32Array, meanValues } = this.formatAudioData(channels);
const rawAudioData = this.floatTo16BitPCM(float32Array);
const monoAudioData = this.floatTo16BitPCM(meanValues);
this.port.postMessage({
event: 'chunk',
data: {
mono: monoAudioData,
raw: rawAudioData,
},
});
}
process(inputList, outputList, parameters) {
// Copy input to output (e.g. speakers)
// Note that this creates choppy sounds with Mac products
const sourceLimit = Math.min(inputList.length, outputList.length);
for (let inputNum = 0; inputNum < sourceLimit; inputNum++) {
const input = inputList[inputNum];
const output = outputList[inputNum];
const channelCount = Math.min(input.length, output.length);
for (let channelNum = 0; channelNum < channelCount; channelNum++) {
input[channelNum].forEach((sample, i) => {
output[channelNum][i] = sample;
});
}
}
const inputs = inputList[0];
// There's latency at the beginning of a stream before recording starts
// Make sure we actually receive audio data before we start storing chunks
let sliceIndex = 0;
if (!this.foundAudio) {
for (const channel of inputs) {
sliceIndex = 0; // reset for each channel
if (this.foundAudio) {
break;
}
if (channel) {
for (const value of channel) {
if (value !== 0) {
// find only one non-zero entry in any channel
this.foundAudio = true;
break;
} else {
sliceIndex++;
}
}
}
}
}
if (inputs && inputs[0] && this.foundAudio && this.recording) {
// We need to copy the TypedArray, because the \`process\`
// internals will reuse the same buffer to hold each input
const chunk = inputs.map((input) => input.slice(sliceIndex));
this.chunks.push(chunk);
this.sendChunk(chunk);
}
return true;
}
}
registerProcessor('audio_processor', AudioProcessor);
`;
const script = new Blob([AudioProcessorWorklet], {
type: 'application/javascript'
});
const src = URL.createObjectURL(script);
export const AudioProcessorSrc = src;
@@ -0,0 +1,96 @@
export const StreamProcessorWorklet = `
class StreamProcessor extends AudioWorkletProcessor {
constructor() {
super();
this.hasStarted = false;
this.hasInterrupted = false;
this.outputBuffers = [];
this.bufferLength = 128;
this.write = { buffer: new Float32Array(this.bufferLength), trackId: null };
this.writeOffset = 0;
this.trackSampleOffsets = {};
this.port.onmessage = (event) => {
if (event.data) {
const payload = event.data;
if (payload.event === 'write') {
const int16Array = payload.buffer;
const float32Array = new Float32Array(int16Array.length);
for (let i = 0; i < int16Array.length; i++) {
float32Array[i] = int16Array[i] / 0x8000; // Convert Int16 to Float32
}
this.writeData(float32Array, payload.trackId);
} else if (
payload.event === 'offset' ||
payload.event === 'interrupt'
) {
const requestId = payload.requestId;
const trackId = this.write.trackId;
const offset = this.trackSampleOffsets[trackId] || 0;
this.port.postMessage({
event: 'offset',
requestId,
trackId,
offset,
});
if (payload.event === 'interrupt') {
this.hasInterrupted = true;
}
} else {
throw new Error(\`Unhandled event "\${payload.event}"\`);
}
}
};
}
writeData(float32Array, trackId = null) {
let { buffer } = this.write;
let offset = this.writeOffset;
for (let i = 0; i < float32Array.length; i++) {
buffer[offset++] = float32Array[i];
if (offset >= buffer.length) {
this.outputBuffers.push(this.write);
this.write = { buffer: new Float32Array(this.bufferLength), trackId };
buffer = this.write.buffer;
offset = 0;
}
}
this.writeOffset = offset;
return true;
}
process(inputs, outputs, parameters) {
const output = outputs[0];
const outputChannelData = output[0];
const outputBuffers = this.outputBuffers;
if (this.hasInterrupted) {
this.port.postMessage({ event: 'stop' });
return false;
} else if (outputBuffers.length) {
this.hasStarted = true;
const { buffer, trackId } = outputBuffers.shift();
for (let i = 0; i < outputChannelData.length; i++) {
outputChannelData[i] = buffer[i] || 0;
}
if (trackId) {
this.trackSampleOffsets[trackId] =
this.trackSampleOffsets[trackId] || 0;
this.trackSampleOffsets[trackId] += buffer.length;
}
return true;
} else if (this.hasStarted) {
this.port.postMessage({ event: 'stop' });
return false;
} else {
return true;
}
}
}
registerProcessor('stream_processor', StreamProcessor);
`;
const script = new Blob([StreamProcessorWorklet], {
type: 'application/javascript'
});
const src = URL.createObjectURL(script);
export const StreamProcessorSrc = src;
+7
View File
@@ -0,0 +1,7 @@
{
// This is for tsup build since it's not supporting composite
"extends": "./tsconfig.json",
"compilerOptions": {
"composite": false
}
}
+39
View File
@@ -0,0 +1,39 @@
{
"compilerOptions": {
// THIS MUST BE AT ROOT, if you set baseurl in sub-package it breaks intellisense jump to
"composite": true,
"baseUrl": ".",
"rootDir": ".",
"outDir": "dist",
"importHelpers": true,
"allowJs": false,
"allowSyntheticDefaultImports": true,
"forceConsistentCasingInFileNames": true,
"declaration": true,
"downlevelIteration": true,
"strict": true,
"esModuleInterop": true,
"jsx": "react-jsx",
"module": "system",
"moduleResolution": "node",
"noEmitOnError": false,
"noImplicitAny": false,
"noImplicitReturns": false,
"noUnusedLocals": false,
"noUnusedParameters": false,
"preserveConstEnums": true,
"removeComments": true,
"skipLibCheck": true,
"sourceMap": true,
"strictNullChecks": true,
"target": "es5",
"types": ["node", "react"],
"lib": ["DOM", "DOM.Iterable", "ESNext"],
"paths": {
"src/*": ["./src/*"]
}
},
"exclude": ["**/test", "**/dist", "**/__tests__"],
"include": ["src/**/*"],
"types": ["@testing-library/jest-dom", "node"]
}