@@ -7,12 +7,14 @@ import { config } from "../Core/Config";
77import { ChannelSessionRow , getChannelSession } from "../Core/Database" ;
88import { channelManager } from "../Channel/ChannelManager" ;
99import { createProcessAIHandler } from "../Processing/createProcessAIHandler" ;
10- import { classifyIntent } from "../Processing/classifyIntent" ;
10+ import { MiddlewarePipeline } from "../Middleware/MiddlewarePipeline" ;
11+ import { intentFilterMiddleware } from "../Middleware/intentFilter" ;
12+ import type { MessageContext } from "../Middleware/types" ;
1113
1214import { getBuiltInCommands } from "./BuiltInCommands" ;
1315import { WebSocketSessionHandler } from "../Channel/web/WebSocketSessionHandler" ;
1416
15- interface ChannelRouteArgs extends ChannelMessageArgs {
17+ export interface ChannelRouteArgs extends ChannelMessageArgs {
1618 channelType : string ;
1719 channelId : string ;
1820 dbSessionId : number ;
@@ -148,9 +150,11 @@ function mergeMessageContents(items: { query: MessageContent }[]): MessageConten
148150
149151export class SbotSessionManager extends SessionManager {
150152 private mergeBuffers = new Map < string , MergeBufferEntry > ( ) ;
153+ readonly messagePipeline = new MiddlewarePipeline < MessageContext > ( ) ;
151154
152155 constructor ( ) {
153156 super ( ) ;
157+ this . messagePipeline . use ( intentFilterMiddleware ) ;
154158 }
155159
156160 protected createSession ( threadId : string ) : SessionService {
@@ -216,20 +220,12 @@ export class SbotSessionManager extends SessionManager {
216220 }
217221
218222 private async dispatchToSession ( threadId : string , query : MessageContent , args : ChannelRouteArgs ) : Promise < void > {
219- if ( ! await this . passIntentFilter ( query , args , threadId ) ) return ;
220- const session = this . getOrCreate ( threadId ) ;
221- await session . onReceiveMessage ( query , args ) ;
222- }
223-
224- private async passIntentFilter ( query : MessageContent , args : ChannelRouteArgs , threadId : string ) : Promise < boolean > {
225- if ( args ?. mentionBot ) return true ;
226- const dbSession = await getChannelSession ( args ?. dbSessionId ) ;
227- const channel = args . channelId ? config . getChannel ( args . channelId ) : undefined ;
228- const intentModel = dbSession ?. intentModel != null ? dbSession . intentModel : channel ?. intentModel ;
229- if ( ! intentModel ) return true ;
230- const intentPrompt = dbSession ?. intentPrompt != null ? dbSession . intentPrompt : ( channel ?. intentPrompt ?? null ) ;
231- const intentThreshold = dbSession ?. intentThreshold != null ? dbSession . intentThreshold : ( channel ?. intentThreshold ?? 0.7 ) ;
232- return classifyIntent ( query , intentModel , intentPrompt , intentThreshold , threadId ) ;
223+ const ctx : MessageContext = { query, args, threadId, filtered : false } ;
224+ await this . messagePipeline . execute ( ctx , async ( c ) => {
225+ if ( c . filtered ) return ;
226+ const session = this . getOrCreate ( c . threadId ) ;
227+ await session . onReceiveMessage ( c . query , c . args ) ;
228+ } ) ;
233229 }
234230
235231 async onReceiveWebMessage ( threadId : string , query : MessageContent , sessionId : string , dbSessionId : number ) : Promise < void > {
0 commit comments