RAG/utils/metadata_filter.py

161 lines
5.2 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

#!/usr/bin/env python3
"""
Metadata过滤工具
"""
import re
import json
from typing import Dict, Any, Optional
from llama_index.llms.ollama import Ollama
from config import settings
class MetadataFilter:
"""
Metadata过滤工具类
"""
@staticmethod
def extract_filters(query: str) -> Dict[str, Any]:
"""
从查询字符串中提取metadata过滤条件
Args:
query: 用户查询字符串
Returns:
包含过滤条件的字典,格式为 {metadata_key: filter_value}
"""
filters = {}
# 提取路径信息,例如:"在 data_structures 目录下"
path_patterns = [
r'\s*([^\s]+)\s*目录下',
r'\s*([^\s]+)\s*文件夹下',
r'路径\s*([^\s]+)'
]
for pattern in path_patterns:
match = re.search(pattern, query)
if match:
path = match.group(1)
filters['file_path'] = path
break
# 提取其他metadata字段的过滤条件
# 例如:"类型为python的文件"、"语言为java的代码"
type_patterns = [
r'类型为\s*([^\s]+)',
r'语言为\s*([^\s]+)'
]
for pattern in type_patterns:
match = re.search(pattern, query)
if match:
lang = match.group(1)
filters['lang'] = lang
break
# 可以根据需要添加更多的过滤条件提取规则
return filters
@staticmethod
def apply_filter(metadata: Dict[str, Any], filters: Dict[str, Any]) -> bool:
"""
应用过滤条件到metadata
Args:
metadata: 文档的metadata
filters: 过滤条件
Returns:
bool: 如果metadata符合过滤条件返回True否则返回False
"""
if not filters:
return True
for key, value in filters.items():
if key not in metadata:
return False
metadata_value = metadata[key]
if isinstance(metadata_value, str):
if value not in metadata_value:
return False
else:
if metadata_value != value:
return False
return True
@staticmethod
def extract_filters_with_llm(query: str, llm: Optional[Ollama] = None) -> Dict[str, Any]:
"""
使用LLM从查询字符串中提取metadata过滤条件
Args:
query: 用户查询字符串
llm: LLM实例如果为None则创建新实例
Returns:
包含过滤条件的字典,格式为 {metadata_key: filter_value}
"""
if not llm:
llm = Ollama(
model=settings.OLLAMA_MODEL,
base_url=settings.OLLAMA_BASE_URL,
temperature=0.1,
request_timeout=30.0
)
# 构建提示词
prompt = f"""
请从以下用户查询中提取metadata过滤条件
{query}
请严格按照以下规则提取:
1. 只提取明确在查询中提到的过滤条件,不要进行任何推测
2. 只提取与以下metadata key一致的过滤条件
- func_id: 函数ID
- func_name: 函数名
- class_name: 类名
- file_path: 文件路径
- lang: 编程语言
- params: 参数数量
- return_type: 返回类型
- docstring: 文档字符串
- start_line: 开始行号
- end_line: 结束行号
- repo_id: 仓库ID
- branch: 分支名
- func_body: 函数体
3. 只有当查询中明确提到某个key的值时才将其包含在结果中
4. 例如:对于查询"在 algorithms 目录下用java实现的排序算法"
只提取 {{"file_path": "algorithms", "lang": "java"}}不要提取其他任何key
请以JSON格式返回提取的过滤条件格式为
{{"file_path": "value", "lang": "value", ...}}
如果没有找到任何过滤条件,请返回空对象:{{}}
"""
# 调用LLM
response = llm.complete(prompt)
# 解析响应
try:
filters = json.loads(response.text)
if isinstance(filters, dict):
# 验证并过滤结果确保只包含有效的metadata key
valid_keys = ['func_id', 'func_name', 'class_name', 'file_path', 'lang', 'params', 'return_type', 'docstring', 'start_line', 'end_line', 'repo_id', 'branch', 'func_body']
filtered_filters = {}
for key, value in filters.items():
if key in valid_keys and value:
filtered_filters[key] = value
return filtered_filters
else:
return {}
except json.JSONDecodeError:
# 如果解析失败,回退到正则表达式提取
return MetadataFilter.extract_filters(query)