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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 16 additions & 0 deletions prisma/schema.prisma
Original file line number Diff line number Diff line change
Expand Up @@ -140,6 +140,22 @@ model Message {
@@index([conversationId])
}

model ContextDocument {
id String @id @default(uuid())
title String
category String
content String
tags String // Stored as JSON string
sourceUrl String?
isActive Boolean @default(true)
createdBy String?
createdAt DateTime @default(now())
updatedAt DateTime @updatedAt

@@index([category])
@@index([isActive])
}

model AiUsageMetric {
id String @id @default(uuid())
userId String?
Expand Down
4 changes: 3 additions & 1 deletion src/ai-assistant/ai-assistant.controller.ts
Original file line number Diff line number Diff line change
@@ -1,9 +1,10 @@
import { Controller, Get, Post, Body, Param, Delete, UseGuards, Req } from '@nestjs/common';
import { ApiTags, ApiOperation, ApiResponse, ApiBearerAuth } from '@nestjs/swagger';
import { AiAssistantService } from './ai-assistant.service';
import { AiAssistantService } from './services/ai-assistant.service';
import { CreateConversationDto, SendMessageDto } from './dto/ai-assistant.dto';
import { JwtAuthGuard } from '../auth/jwt-auth.guard';
import { CurrentUser } from '../common/decorators/current-user.decorator';
import { ThrottleByWallet } from '../common/decorators/throttle-by-wallet.decorator';

@ApiTags('AI Assistant')
@ApiBearerAuth()
Expand Down Expand Up @@ -36,6 +37,7 @@ export class AiAssistantController {
@Post(':id/messages')
@ApiOperation({ summary: 'Send a message to the AI Assistant' })
@ApiResponse({ status: 201, description: 'AI Assistant response.' })
@ThrottleByWallet('ai')
async sendMessage(
@CurrentUser() user: any,
@Param('id') conversationId: string,
Expand Down
15 changes: 9 additions & 6 deletions src/ai-assistant/ai-assistant.module.ts
Original file line number Diff line number Diff line change
@@ -1,15 +1,18 @@
import { Module } from '@nestjs/common';
import { TypeOrmModule } from '@nestjs/typeorm';
import { AiAssistantController } from './ai-assistant.controller';
import { AiAssistantService } from './ai-assistant.service';
import { LlmProviderService } from './llm-provider.service';
import { RagService } from './rag.service';
import { AiAssistantService } from './services/ai-assistant.service';
import { LlmProviderService } from './services/llm-provider.service';
import { RagService } from './services/rag.service';
import { SafetyGuardrailService } from './services/safety-guardrail.service';
import { ContextDocument } from './entities/context-document.entity';
import { PrismaModule } from '../prisma/prisma.module';
// Note: assuming PrismaModule is exported from '../prisma/prisma.module'
import { RedisModule } from '../redis/redis.module';

@Module({
imports: [PrismaModule],
imports: [PrismaModule, RedisModule, TypeOrmModule.forFeature([ContextDocument])],
controllers: [AiAssistantController],
providers: [AiAssistantService, LlmProviderService, RagService],
providers: [AiAssistantService, LlmProviderService, RagService, SafetyGuardrailService],
exports: [AiAssistantService],
})
export class AiAssistantModule {}
28 changes: 0 additions & 28 deletions src/ai-assistant/rag.service.ts

This file was deleted.

Original file line number Diff line number Diff line change
@@ -1,8 +1,9 @@
import { Test, TestingModule } from '@nestjs/testing';
import { AiAssistantService } from './ai-assistant.service';
import { PrismaService } from '../prisma/prisma.service';
import { PrismaService } from '../../prisma/prisma.service';
import { LlmProviderService } from './llm-provider.service';
import { RagService } from './rag.service';
import { SafetyGuardrailService } from './safety-guardrail.service';

describe('AiAssistantService', () => {
let service: AiAssistantService;
Expand Down Expand Up @@ -36,7 +37,11 @@ describe('AiAssistantService', () => {
};

const mockRagService = {
retrieveContext: jest.fn().mockResolvedValue('mock context'),
retrieveContext: jest.fn().mockResolvedValue({ context: 'mock context', citations: [] }),
};

const mockSafetyGuardrail = {
checkContent: jest.fn().mockReturnValue({ flagged: false }),
};

const module: TestingModule = await Test.createTestingModule({
Expand All @@ -45,6 +50,7 @@ describe('AiAssistantService', () => {
{ provide: PrismaService, useValue: mockPrismaService },
{ provide: LlmProviderService, useValue: mockLlmProvider },
{ provide: RagService, useValue: mockRagService },
{ provide: SafetyGuardrailService, useValue: mockSafetyGuardrail },
],
}).compile();

Expand Down
Original file line number Diff line number Diff line change
@@ -1,8 +1,9 @@
import { Injectable, NotFoundException, Logger, ForbiddenException } from '@nestjs/common';
import { PrismaService } from '../prisma/prisma.service';
import { PrismaService } from '../../prisma/prisma.service';
import { LlmProviderService } from './llm-provider.service';
import { RagService } from './rag.service';
import { CreateConversationDto, SendMessageDto } from './dto/ai-assistant.dto';
import { SafetyGuardrailService } from './safety-guardrail.service';
import { CreateConversationDto, SendMessageDto } from '../dto/ai-assistant.dto';

@Injectable()
export class AiAssistantService {
Expand All @@ -12,6 +13,7 @@ export class AiAssistantService {
private prisma: PrismaService,
private llmProvider: LlmProviderService,
private ragService: RagService,
private safetyGuardrail: SafetyGuardrailService,
) {}

async createConversation(userId: string, dto: CreateConversationDto) {
Expand Down Expand Up @@ -62,6 +64,9 @@ export class AiAssistantService {
throw new ForbiddenException('You do not have access to this conversation');
}

// 0. Safety Check
const safetyCheck = this.safetyGuardrail.checkContent(dto.content);

// 1. Save user message
const userMessage = await this.prisma.message.create({
data: {
Expand All @@ -71,6 +76,28 @@ export class AiAssistantService {
},
});

if (safetyCheck.flagged) {
const assistantMessage = await this.prisma.message.create({
data: {
conversationId,
role: 'assistant',
content: 'I cannot answer this request.',
},
});

return {
message: assistantMessage,
metadata: {
provider: 'none',
latencyMs: 0,
tokens: 0,
citations: [],
flagged: true,
flagReason: safetyCheck.reason,
}
};
}

// 2. Retrieve Conversation History
const history = await this.prisma.message.findMany({
where: { conversationId },
Expand All @@ -79,7 +106,7 @@ export class AiAssistantService {
});

// 3. RAG Retrieval
const context = await this.ragService.retrieveContext(dto.content);
const { context, citations } = await this.ragService.retrieveContext(dto.content);

// 4. Construct Prompt Pipeline
const systemPrompt = `You are the TruthBounty AI Assistant. You help contributors navigate the protocol.
Expand Down Expand Up @@ -139,7 +166,7 @@ ${context}
provider: llmResponse.provider,
latencyMs,
tokens: llmResponse.usage?.total_tokens || 0,
citations: ['MOCKED_CITATION_1', 'MOCKED_CITATION_2'] // Placeholder for standardizing API
citations
}
};
}
Expand Down
94 changes: 94 additions & 0 deletions src/ai-assistant/services/llm-provider.service.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,94 @@
import { Injectable, Logger } from '@nestjs/common';
import { ConfigService } from '@nestjs/config';
import OpenAI from 'openai';
import Anthropic from '@anthropic-ai/sdk';

@Injectable()
export class LlmProviderService {
private readonly logger = new Logger(LlmProviderService.name);
private openai: OpenAI | null = null;
private anthropic: Anthropic | null = null;
private defaultProvider: 'openai' | 'anthropic';

constructor(private configService: ConfigService) {
const openaiKey = this.configService.get<string>('OPENAI_API_KEY');
if (openaiKey) {
this.openai = new OpenAI({ apiKey: openaiKey });
}

const anthropicKey = this.configService.get<string>('ANTHROPIC_API_KEY');
if (anthropicKey) {
this.anthropic = new Anthropic({ apiKey: anthropicKey });
}

this.defaultProvider = this.configService.get<'openai' | 'anthropic'>('DEFAULT_LLM_PROVIDER') || 'openai';
}

async generateEmbedding(text: string): Promise<number[]> {
if (this.openai) {
const response = await this.openai.embeddings.create({
model: 'text-embedding-3-small',
input: text,
});
return response.data[0].embedding;
}
this.logger.warn('OpenAI not configured, returning mock embedding.');
return new Array(1536).fill(0.1);
}

async generateResponse(
messages: { role: 'user' | 'assistant' | 'system'; content: string }[],
options?: { provider?: 'openai' | 'anthropic' }
): Promise<{ content: string; usage: any; provider: string; model: string }> {
const provider = options?.provider || this.defaultProvider;

if (provider === 'openai' && this.openai) {
const model = 'gpt-4o-mini';
const response = await this.openai.chat.completions.create({
model,
messages: messages.map(m => ({ role: m.role, content: m.content })),
});
return {
content: response.choices[0].message.content || '',
usage: response.usage,
provider: 'openai',
model,
};
} else if (provider === 'anthropic' && this.anthropic) {
const model = 'claude-3-haiku-20240307';
const systemMessage = messages.find(m => m.role === 'system')?.content;
const otherMessages = messages.filter(m => m.role !== 'system').map(m => ({
role: m.role === 'assistant' ? 'assistant' as const : 'user' as const,
content: m.content
}));

const response = await this.anthropic.messages.create({
model,
max_tokens: 1024,
system: systemMessage,
messages: otherMessages,
});

const content = response.content[0].type === 'text' ? response.content[0].text : '';
return {
content,
usage: {
prompt_tokens: response.usage.input_tokens,
completion_tokens: response.usage.output_tokens,
total_tokens: response.usage.input_tokens + response.usage.output_tokens,
},
provider: 'anthropic',
model,
};
}

// Mock fallback if keys not configured
this.logger.warn(`No valid LLM provider configured for ${provider}, using mock response.`);
return {
content: `This is a mock response from the AI Assistant because the API keys for ${provider} are not configured. You said: ${messages[messages.length - 1]?.content}`,
usage: { prompt_tokens: 10, completion_tokens: 20, total_tokens: 30 },
provider: 'mock',
model: 'mock-model',
};
}
}
24 changes: 24 additions & 0 deletions src/ai-assistant/services/rag.service.spec.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,24 @@
import { Test, TestingModule } from '@nestjs/testing';
import { RagService } from './rag.service';
import { PrismaService } from '../../prisma/prisma.service';
import { LlmProviderService } from './llm-provider.service';

describe('RagService', () => {
let service: RagService;

beforeEach(async () => {
const module: TestingModule = await Test.createTestingModule({
providers: [
RagService,
{ provide: PrismaService, useValue: { contextDocument: { findMany: jest.fn().mockResolvedValue([]) } } },
{ provide: LlmProviderService, useValue: {} },
],
}).compile();

service = module.get<RagService>(RagService);
});

it('should be defined', () => {
expect(service).toBeDefined();
});
});
64 changes: 64 additions & 0 deletions src/ai-assistant/services/rag.service.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,64 @@
import { Injectable, Logger } from '@nestjs/common';
import { InjectRepository } from '@nestjs/typeorm';
import { Repository } from 'typeorm';
import { ContextDocument } from '../entities/context-document.entity';
import { RedisService } from '../../redis/redis.service';

@Injectable()
export class RagService {
private readonly logger = new Logger(RagService.name);

constructor(
@InjectRepository(ContextDocument)
private readonly contextDocumentRepository: Repository<ContextDocument>,
private redisService: RedisService,
) {}

async retrieveContext(query: string): Promise<{ context: string; citations: string[] }> {
const cacheKey = `rag_context:${query.trim().toLowerCase()}`;
const cached = await this.redisService.get(cacheKey);
if (cached) {
this.logger.debug(`Cache hit for query: ${query}`);
return JSON.parse(cached);
}

this.logger.debug(`Retrieving context for query: ${query}`);

// 1. Fetch all active documents
const documents = await this.contextDocumentRepository.find({
where: { isActive: true },
});

if (documents.length === 0) {
return { context: 'No protocol documentation found.', citations: [] };
}

// 2. Simple keyword-based ranking for now as a fallback
const relevantDocs = documents
.map(doc => ({
...doc,
score: this.calculateRelevance(query, doc.content + ' ' + doc.title)
}))
.sort((a, b) => b.score - a.score)
.slice(0, 3); // Take top 3

const result = {
context: relevantDocs.map(doc => `[${doc.title}]: ${doc.content}`).join('\n\n'),
citations: relevantDocs.map(doc => doc.title)
};

await this.redisService.set(cacheKey, JSON.stringify(result), 3600); // 1 hour cache
return result;
}

private calculateRelevance(query: string, content: string): number {
const queryTerms = query.toLowerCase().split(/\s+/);
let score = 0;
queryTerms.forEach(term => {
if (content.toLowerCase().includes(term)) {
score += 1;
}
});
return score;
}
}
Loading
Loading