Repository navigation
Expand file tree
/
Copy pathembedding_service.py
More file actions
129 lines (103 loc) · 4.71 KB
/
Copy pathembedding_service.py
File metadata and controls
129 lines (103 loc) · 4.71 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
import openai
import numpy as np
from typing import List, Dict, Optional
from config import Config
import logging
import time
logger = logging.getLogger(__name__)
class EmbeddingService:
def __init__(self, api_key: Optional[str] = None, model: Optional[str] = None):
self.api_key = api_key or Config.OPENAI_API_KEY
self.model = model or Config.OPENAI_MODEL
if not self.api_key:
raise ValueError("OpenAI API key is required")
# Create OpenAI client without setting global api_key
self.client = openai.OpenAI(api_key=self.api_key)
def get_embedding(self, text: str) -> Optional[List[float]]:
"""Get embedding for a single text"""
try:
# Clean and prepare text
text = text.strip()
if not text:
return None
# Truncate if too long (OpenAI has limits)
max_tokens = 8191 # OpenAI's limit for text-embedding-ada-002
if len(text) > max_tokens * 4: # Rough estimate: 4 chars per token
text = text[:max_tokens * 4]
response = self.client.embeddings.create(
model=self.model,
input=text
)
return response.data[0].embedding
except Exception as e:
logger.error(f"Error getting embedding: {e}")
return None
def get_embeddings_batch(self, texts: List[str], batch_size: int = 100) -> List[Optional[List[float]]]:
"""Get embeddings for multiple texts in batches"""
embeddings = []
for i in range(0, len(texts), batch_size):
batch = texts[i:i + batch_size]
logger.info(f"Processing embedding batch {i//batch_size + 1}/{(len(texts) + batch_size - 1)//batch_size}")
try:
# Clean and prepare batch
cleaned_batch = []
for text in batch:
text = text.strip()
if text:
# Truncate if too long
max_tokens = 8191
if len(text) > max_tokens * 4:
text = text[:max_tokens * 4]
cleaned_batch.append(text)
else:
cleaned_batch.append("")
if not any(cleaned_batch):
embeddings.extend([None] * len(batch))
continue
response = self.client.embeddings.create(
model=self.model,
input=cleaned_batch
)
# Map embeddings back to original order
batch_embeddings = []
response_index = 0
for text in cleaned_batch:
if text:
batch_embeddings.append(response.data[response_index].embedding)
response_index += 1
else:
batch_embeddings.append(None)
embeddings.extend(batch_embeddings)
# Rate limiting - be nice to OpenAI
time.sleep(0.1)
except Exception as e:
logger.error(f"Error processing embedding batch: {e}")
embeddings.extend([None] * len(batch))
return embeddings
def process_chunks(self, chunks: List[Dict]) -> List[Dict]:
"""Process chunks and add embeddings"""
if not chunks:
return []
# Extract texts for embedding
texts = [chunk['content'] for chunk in chunks]
# Get embeddings
embeddings = self.get_embeddings_batch(texts)
# Add embeddings to chunks
processed_chunks = []
for chunk, embedding in zip(chunks, embeddings):
if embedding is not None:
chunk_copy = chunk.copy()
chunk_copy['embedding'] = embedding
processed_chunks.append(chunk_copy)
else:
logger.warning(f"Failed to get embedding for chunk: {chunk.get('file_path', 'unknown')}")
logger.info(f"Successfully processed {len(processed_chunks)}/{len(chunks)} chunks with embeddings")
return processed_chunks
def test_connection(self) -> bool:
"""Test OpenAI API connection"""
try:
test_embedding = self.get_embedding("test")
return test_embedding is not None
except Exception as e:
logger.error(f"OpenAI API connection test failed: {e}")
return False