feat(adapters): add SSH and database adapters
This commit is contained in:
@@ -0,0 +1,92 @@
|
||||
import os
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
from typing import List, AsyncIterator, Dict, Any
|
||||
import paramiko
|
||||
from .base import SourceAdapter, FileInfo
|
||||
|
||||
|
||||
class SSHAdapter(SourceAdapter):
|
||||
def __init__(self, config: Dict[str, Any]):
|
||||
super().__init__(config)
|
||||
self.client = None
|
||||
self.sftp = None
|
||||
|
||||
async def connect(self) -> None:
|
||||
self.client = paramiko.SSHClient()
|
||||
self.client.set_missing_host_key_policy(paramiko.AutoAddPolicy())
|
||||
|
||||
host = self.config.get("host", "localhost")
|
||||
port = self.config.get("port", 22)
|
||||
username = self.config.get("username")
|
||||
password = self.config.get("password")
|
||||
key_path = self.config.get("key_path")
|
||||
|
||||
connect_kwargs = {
|
||||
"hostname": host,
|
||||
"port": port,
|
||||
"username": username,
|
||||
}
|
||||
|
||||
if password:
|
||||
connect_kwargs["password"] = password
|
||||
elif key_path and os.path.exists(key_path):
|
||||
connect_kwargs["key_filename"] = key_path
|
||||
|
||||
self.client.connect(**connect_kwargs)
|
||||
self.sftp = self.client.open_sftp()
|
||||
|
||||
async def disconnect(self) -> None:
|
||||
if self.sftp:
|
||||
self.sftp.close()
|
||||
self.sftp = None
|
||||
if self.client:
|
||||
self.client.close()
|
||||
self.client = None
|
||||
|
||||
async def list_files(self, path: str = "") -> List[FileInfo]:
|
||||
remote_path = self.config.get("path", ".")
|
||||
target_path = f"{remote_path}/{path}" if path else remote_path
|
||||
|
||||
files = []
|
||||
exclude_patterns = self.config.get("exclude", [])
|
||||
|
||||
try:
|
||||
for entry in self.sftp.listdir_attr(target_path):
|
||||
entry_path = f"{target_path}/{entry.filename}"
|
||||
rel_path = entry_path.replace(remote_path + "/", "", 1) if remote_path != "." else entry_path
|
||||
|
||||
if any(pattern in rel_path for pattern in exclude_patterns):
|
||||
continue
|
||||
|
||||
is_dir = entry.st_mode & 0o40000 == 0o40000 if hasattr(entry, 'st_mode') else False
|
||||
|
||||
files.append(FileInfo(
|
||||
path=rel_path,
|
||||
size=entry.st_size,
|
||||
modified_time=entry.st_mtime,
|
||||
is_directory=is_dir
|
||||
))
|
||||
except IOError:
|
||||
pass
|
||||
|
||||
return files
|
||||
|
||||
async def read_file(self, path: str) -> AsyncIterator[bytes]:
|
||||
remote_path = self.config.get("path", ".")
|
||||
file_path = f"{remote_path}/{path}" if not path.startswith("/") else path
|
||||
|
||||
with tempfile.NamedTemporaryFile(delete=False) as tmp:
|
||||
tmp_path = tmp.name
|
||||
|
||||
try:
|
||||
self.sftp.get(file_path, tmp_path)
|
||||
with open(tmp_path, "rb") as f:
|
||||
while chunk := f.read(8192):
|
||||
yield chunk
|
||||
finally:
|
||||
if os.path.exists(tmp_path):
|
||||
os.unlink(tmp_path)
|
||||
|
||||
async def get_database_dump(self, config: Dict[str, Any]) -> AsyncIterator[bytes]:
|
||||
raise NotImplementedError("SSH adapter does not support database dumps directly")
|
||||
Reference in New Issue
Block a user