adk-typescript
Version:
TypeScript port of Google's Agent Development Kit (ADK)
400 lines (399 loc) • 15 kB
JavaScript
;
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;