UNPKG

adk-typescript

Version:

TypeScript port of Google's Agent Development Kit (ADK)

400 lines (399 loc) 15 kB
"use strict"; Object.defineProperty(exports, "__esModule", { value: true }); exports.VertexAiSessionService = void 0; const BaseSessionService_1 = require("./BaseSessionService"); const sessionUtils_1 = require("./sessionUtils"); /** * Connects to the managed Vertex AI Session Service */ class VertexAiSessionService extends BaseSessionService_1.BaseSessionService { /** * Creates a new VertexAiSessionService * * @param project - The Google Cloud project ID * @param location - The Google Cloud location/region * @param client - Optional GenAi client to use for API calls */ constructor(options) { super(); this.project = options.project || ''; this.location = options.location || ''; // In a real implementation, we would import and use the genai client // For now, we'll just use the provided client or create a stub this.apiClient = options.client?._apiClient || { request: async (options) => { throw new Error('No Vertex AI client provided'); } }; } /** * Creates a new session in Vertex AI */ async createSession(options) { const { appName, userId, state } = options; const reasoningEngineId = this.parseReasoningEngineId(appName); const sessionJsonDict = { user_id: userId }; if (state) { sessionJsonDict.session_state = state; } // Create the session in Vertex AI const apiResponse = await this.apiClient.request({ httpMethod: 'POST', path: `reasoningEngines/${reasoningEngineId}/sessions`, requestDict: sessionJsonDict }); console.log('Create Session response', apiResponse); // Extract session ID and operation ID from the response const sessionId = apiResponse.name.split('/').slice(-3)[0]; const operationId = apiResponse.name.split('/').slice(-1)[0]; // Poll for operation completion let maxRetryAttempt = 5; while (maxRetryAttempt >= 0) { const lroResponse = await this.apiClient.request({ httpMethod: 'GET', path: `operations/${operationId}`, requestDict: {} }); if (lroResponse.done) { break; } await new Promise(resolve => setTimeout(resolve, 1000)); maxRetryAttempt--; } // Get the session resource const getSessionApiResponse = await this.apiClient.request({ httpMethod: 'GET', path: `reasoningEngines/${reasoningEngineId}/sessions/${sessionId}`, requestDict: {} }); // Parse the update timestamp const updateTimestamp = new Date(getSessionApiResponse.updateTime).getTime() / 1000; // Create and return the session const session = { id: sessionId, appName: appName, userId: userId, state: getSessionApiResponse.sessionState || {}, events: [] }; return session; } /** * Gets a session by its ID */ async getSession(options) { const { appName, userId, sessionId } = options; const reasoningEngineId = this.parseReasoningEngineId(appName); try { // Get session resource const getSessionApiResponse = await this.apiClient.request({ httpMethod: 'GET', path: `reasoningEngines/${reasoningEngineId}/sessions/${sessionId}`, requestDict: {} }); // Parse the update timestamp const updateTimestamp = new Date(getSessionApiResponse.updateTime).getTime() / 1000; // Create the session const session = { id: sessionId, appName: appName, userId: userId, state: getSessionApiResponse.sessionState || {}, events: [] }; // Get the session events const listEventsApiResponse = await this.apiClient.request({ httpMethod: 'GET', path: `reasoningEngines/${reasoningEngineId}/sessions/${sessionId}/events`, requestDict: {} }); // Handle empty response case if (listEventsApiResponse.httpHeaders) { return session; } // Convert API events to Event objects session.events = listEventsApiResponse.sessionEvents .map((event) => this.fromApiEvent(event)) .filter((event) => { // Filter events by timestamp return event.timestamp !== undefined && event.timestamp <= updateTimestamp; }) .sort((a, b) => { // Sort events by timestamp return (a.timestamp || 0) - (b.timestamp || 0); }); return session; } catch (error) { console.error('Error getting session:', error); return null; } } /** * Lists all sessions for a user in an app */ async listSessions(options) { const { appName, userId } = options; const reasoningEngineId = this.parseReasoningEngineId(appName); try { const apiResponse = await this.apiClient.request({ httpMethod: 'GET', path: `reasoningEngines/${reasoningEngineId}/sessions?filter=user_id=${userId}`, requestDict: {} }); // Handle empty response case if (apiResponse.httpHeaders) { return { sessions: [] }; } // Convert API sessions to Session objects const sessions = apiResponse.sessions.map((apiSession) => { return { id: apiSession.name.split('/').slice(-1)[0], appName: appName, userId: userId, state: {}, // Don't load full state for listing events: [] // Don't load events for listing }; }); return { sessions }; } catch (error) { console.error('Error listing sessions:', error); return { sessions: [] }; } } /** * Deletes a session */ async deleteSession(options) { const { appName, sessionId } = options; const reasoningEngineId = this.parseReasoningEngineId(appName); try { await this.apiClient.request({ httpMethod: 'DELETE', path: `reasoningEngines/${reasoningEngineId}/sessions/${sessionId}`, requestDict: {} }); } catch (error) { console.error('Error deleting session:', error); } } /** * Lists events in a session */ async listEvents(options) { const { appName, sessionId } = options; const reasoningEngineId = this.parseReasoningEngineId(appName); try { const apiResponse = await this.apiClient.request({ httpMethod: 'GET', path: `reasoningEngines/${reasoningEngineId}/sessions/${sessionId}/events`, requestDict: {} }); console.log('List events response', apiResponse); // Handle empty response case if (apiResponse.httpHeaders) { return { events: [] }; } // Convert API events to Event objects const events = apiResponse.sessionEvents.map((event) => this.fromApiEvent(event)); return { events }; } catch (error) { console.error('Error listing events:', error); return { events: [] }; } } /** * Appends an event to a session */ async appendEvent(options) { const { session, event } = options; // Update the in-memory session super.appendEvent(options); // Update the session in Vertex AI const reasoningEngineId = this.parseReasoningEngineId(session.appName); try { await this.apiClient.request({ httpMethod: 'POST', path: `reasoningEngines/${reasoningEngineId}/sessions/${session.id}:appendEvent`, requestDict: this.convertEventToJson(event) }); } catch (error) { console.error('Error appending event:', error); } } /** * Updates a session's state. */ async updateSessionState(appName, userId, sessionId, stateDelta) { const reasoningEngineId = this.parseReasoningEngineId(appName); try { // Get the current session const session = await this.getSession({ appName, userId, sessionId }); if (!session) { throw new Error(`Session ${sessionId} not found for user ${userId} in app ${appName}`); } // Update the session state Object.assign(session.state, stateDelta); // Update the session in Vertex AI await this.apiClient.request({ httpMethod: 'PATCH', path: `reasoningEngines/${reasoningEngineId}/sessions/${sessionId}`, requestDict: { session_state: session.state } }); return session; } catch (error) { console.error('Error updating session state:', error); throw error; } } /** * Parses a reasoning engine ID from an app name */ parseReasoningEngineId(appName) { // If app name is just digits, assume it's already a reasoning engine ID if (/^\d+$/.test(appName)) { return appName; } // Check if app name matches the expected format const pattern = /^projects\/([a-zA-Z0-9-_]+)\/locations\/([a-zA-Z0-9-_]+)\/reasoningEngines\/(\d+)$/; const match = appName.match(pattern); if (!match) { throw new Error(`App name ${appName} is not valid. It should either be the full ` + 'ReasoningEngine resource name, or the reasoning engine id.'); } // Return the reasoning engine ID return match[3]; } /** * Converts an Event object to a JSON object for the API */ convertEventToJson(event) { const metadataJson = { partial: event.partial, turn_complete: event.turnComplete, interrupted: event.interrupted, branch: event.branch, long_running_tool_ids: event.longRunningToolIds ? Array.from(event.longRunningToolIds) : undefined }; if (event.groundingMetadata) { metadataJson.grounding_metadata = event.groundingMetadata; } const eventJson = { author: event.author, invocation_id: event.invocationId, timestamp: { seconds: Math.floor(event.timestamp || Date.now() / 1000), nanos: Math.floor(((event.timestamp || Date.now() / 1000) % 1) * 1000000000) }, error_code: event.errorCode, error_message: event.errorMessage, event_metadata: metadataJson }; if (event.actions) { const actionsJson = { skip_summarization: event.actions.skipSummarization, state_delta: event.actions.stateDelta, artifact_delta: event.actions.artifactDelta, transfer_agent: event.actions.transferToAgent, escalate: event.actions.escalate, requested_auth_configs: event.actions.requestedAuthConfigs }; eventJson.actions = actionsJson; } if (event.content) { eventJson.content = (0, sessionUtils_1.encodeContent)(event.content); } return eventJson; } /** * Converts a Content object to a JSON object for the API */ convertContentToJson(content) { return { role: content.role, parts: content.parts.map(part => this.convertPartToJson(part)) }; } /** * Converts a Part object to a JSON object for the API */ convertPartToJson(part) { const result = {}; if (part.text !== undefined) { result.text = part.text; } if (part.data !== undefined && part.mimeType !== undefined) { result.inline_data = { data: Buffer.from(part.data).toString('base64'), mime_type: part.mimeType }; } return result; } /** * Converts an API event to an Event object */ fromApiEvent(apiEvent) { // Parse event actions const eventActions = apiEvent.actions ? { skipSummarization: apiEvent.actions.skipSummarization, stateDelta: apiEvent.actions.stateDelta || {}, artifactDelta: apiEvent.actions.artifactDelta || {}, transferToAgent: apiEvent.actions.transferAgent, escalate: apiEvent.actions.escalate, requestedAuthConfigs: apiEvent.actions.requestedAuthConfigs || {} } : undefined; // Create the event const event = { id: apiEvent.name.split('/').slice(-1)[0], invocationId: apiEvent.invocationId, author: apiEvent.author, content: (0, sessionUtils_1.decodeContent)(apiEvent.content), actions: eventActions, timestamp: new Date(apiEvent.timestamp).getTime() / 1000, errorCode: apiEvent.errorCode, errorMessage: apiEvent.errorMessage }; // Parse event metadata if (apiEvent.eventMetadata) { const longRunningToolIdsList = apiEvent.eventMetadata.longRunningToolIds; event.partial = apiEvent.eventMetadata.partial; event.turnComplete = apiEvent.eventMetadata.turnComplete; event.interrupted = apiEvent.eventMetadata.interrupted; event.branch = apiEvent.eventMetadata.branch; event.groundingMetadata = apiEvent.eventMetadata.groundingMetadata; if (longRunningToolIdsList) { event.longRunningToolIds = new Set(longRunningToolIdsList); } } return event; } /** * Parses content from the API */ parseContent(apiContent) { if (!apiContent) { return { role: 'user', parts: [] }; } return (0, sessionUtils_1.decodeContent)(apiContent); } } exports.VertexAiSessionService = VertexAiSessionService;