161 lines
5.2 KiB
Python
161 lines
5.2 KiB
Python
#!/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)
|